RAG优化(续一)
这里是为了续上上一章节的从RAG入门到RAG优化,所以不多废话,直接上手
RAG优化
(1)chunk优化- [YES]
(2)BM25关键检索- [YES]
(3)rerank 重排
其实很好理解,rerank的实现过程就是首先粗略的找出一些相关的,然后再去使用更聪明的模型去进行排名,然后选出最有用的放在前面,为什么需要重排序呢,其实也很好理解,第一次筛是大范围粗略的筛,所以重排就需要更加的精细,使得检索更加准确
这里是AI生成的流程图,以便于理解整个流程

之所以放上流程图是因为加上许多优化,我们的代码已经比较复杂来到了大概480行左右,如果不理清整体的结构,很明显是无法继续加入优化的。所以这里我们需要加入Rerank我们首先应该知道放在哪里,Rerank的目的是为了当我们检索出来,进行一个得分排序,所以肯定是加入检索过程之后,也就是放在hybrid_search之后
#核心代码
#创建排序器
from langchain_community.document_compressors.dashscope_rerank import DashScopeRerank
def create_reranker():
"""
创建排序器,所以直接返回,为了主函数里直接实例化
"""
return DashScopeRerank(
model="..."
api_key="..."
#这里不用谢base_url因为用的是langchain适配的DashScopeRerank,这里查询之后没有OpenAI的接口通用写法,所以只能用这个
#默认base_url为阿里百炼的网站,所以可以不写
top_n=...
#保留前n个
)
def hybrid_search(
user_question,#用户问题
bm25_retriever,
vector_store, #检索器
reranker,#重排器
bm25_k=10,
vector_k=10,
final_k=5,#检索器参数
):
bm25_retriever.k = bm25_k
keyword_docs = bm25_retriever.invoke(user_question)
semantic_docs = vector_store.similarity_search(
user_question,
k=vector_k,
)
candidate_docs = []
seen_keys = set()
for doc in keyword_docs + semantic_docs
unique_key(
doc.metadata.get("source",""),
doc.metadata.get("page",""),
doc.page_content[:80],
)
if unique_key not in seen_keys:
candidate_docs.append(doc)
seen_keys.append(unique_key)
#去重
if not candidate_docs:
return []
reranked_docs = reranker.compress_documents(
documents=candidate_docs,
query=user_question,
)
#重排
return list(reranked_docs)[:final_k]
点击查看完整代码及返回结果
import os
import pickle
import jieba
from dotenv import load_dotenv
from langchain_community.document_loaders import PyPDFLoader
from langchain_community.retrievers import BM25Retriever
from langchain_community.document_compressors.dashscope_rerank import DashScopeRerank
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_chroma import Chroma
from langchain_openai import OpenAIEmbeddings, ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
# =========================
# 1. 全局配置
# =========================
PDF_PATH = r"C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf"
BM25_DOCS_PATH = "./bm25_docs.pkl"
COLLECTION_NAME = "my_rag_collection"
DB_PATH = "./chroma_db"
CHUNK_SIZE = 500
CHUNK_OVERLAP = 100
BM25_K = 10
VECTOR_K = 10
FINAL_K = 5
# =========================
# 2. 文档加载与切分
# =========================
def load_pdf(path):
"""
读取 PDF。
返回 LangChain 的 Document 列表。
"""
loader = PyPDFLoader(path)
return loader.load()
def split_documents(raw_docs, chunk_size=500, chunk_overlap=100):
"""
把大文档切成小块。
chunk 说人话:就是一小段资料。
"""
splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
separators=[
"\n\n",
"\n",
"。",
"!",
"?",
";",
",",
"、",
" ",
"",
],
is_separator_regex=False,
)
chunks = splitter.split_documents(raw_docs)
print(f"切分出 {len(chunks)} 个 chunk")
return chunks
# =========================
# 3. 初始化模型和数据库
# =========================
def create_embedding_model():
"""
创建 embedding 模型。
embedding 说人话:把文字变成数字,方便机器比较相似度。
"""
return OpenAIEmbeddings(
model=os.getenv("DASHSCOPE_EMBEDDING_MODEL"),
api_key=os.getenv("DASHSCOPE_API_KEY"),
base_url=os.getenv("DASHSCOPE_BASE_URL"),
dimensions=1024,
chunk_size=10,
check_embedding_ctx_length=False,
)
def create_vector_store(embedding):
"""
创建 Chroma 向量数据库。
Chroma 说人话:本地知识库。
"""
return Chroma(
collection_name=COLLECTION_NAME,
embedding_function=embedding,
persist_directory=DB_PATH,
)
def create_llm():
"""
创建大模型。
"""
return ChatOpenAI(
model=os.getenv("DASHSCOPE_CHAT_MODEL"),
api_key=os.getenv("DASHSCOPE_API_KEY"),
base_url=os.getenv("DASHSCOPE_BASE_URL"),
)
def create_reranker():
"""
创建 rerank 重排器。
rerank 说人话:把已经找出来的资料重新排队,把最有用的放前面。
"""
return DashScopeRerank(
model=os.getenv("DASHSCOPE_RERANK_MODEL"),
top_n=FINAL_K,
dashscope_api_key=os.getenv("DASHSCOPE_API_KEY"),
)
# =========================
# 4. BM25 相关
# =========================
def chinese_tokenizer(text):
"""
给 BM25 用的中文分词函数。
分词说人话:把一句中文切成一个个词。
"""
return [word.strip() for word in jieba.cut(text) if word.strip()]
def save_bm25_docs(chunks):
"""
保存 chunks,给 BM25 使用。
"""
with open(BM25_DOCS_PATH, "wb") as f:
pickle.dump(chunks, f)
print(f"BM25 文档已保存到 {BM25_DOCS_PATH}")
def load_bm25_retriever(k=5):
"""
加载 BM25 检索器。
BM25 说人话:按关键词找资料。
"""
if not os.path.exists(BM25_DOCS_PATH):
raise FileNotFoundError(
f"找不到 {BM25_DOCS_PATH},请先选择入库模式。"
)
with open(BM25_DOCS_PATH, "rb") as f:
docs = pickle.load(f)
retriever = BM25Retriever.from_documents(
docs,
preprocess_func=chinese_tokenizer,
)
retriever.k = k
return retriever
# =========================
# 5. 入库、删除
# =========================
def ingest_documents(vector_store, chunks):
"""
入库:
1. 存进 Chroma,给向量检索用。
2. 存成 pkl,给 BM25 用。
"""
ids = [f"text_new_chunk_{i}" for i in range(len(chunks))]
vector_store.add_documents(
documents=chunks,
ids=ids,
)
save_bm25_docs(chunks)
print(f"成功入库 {len(chunks)} 个 chunk")
def delete_collection(vector_store):
"""
删除 Chroma collection 和 BM25 文件。
collection 说人话:知识库里的一张表。
"""
vector_store.delete_collection()
print(f"{COLLECTION_NAME} 已删除")
if os.path.exists(BM25_DOCS_PATH):
os.remove(BM25_DOCS_PATH)
print(f"{BM25_DOCS_PATH} 已删除")
# =========================
# 6. 混合检索
# =========================
def hybrid_search(
user_question,
bm25_retriever,
vector_store,
reranker,
bm25_k=10,
vector_k=10,
final_k=5,
):
"""
混合检索 + LangChain rerank:
1. BM25 按关键词找
2. Chroma 按意思找
3. 合并
4. 去重
5. 用 LangChain 封装好的 DashScopeRerank 重排
"""
bm25_retriever.k = bm25_k
# 1. 关键词检索
keyword_docs = bm25_retriever.invoke(user_question)
# 2. 语义检索
semantic_docs = vector_store.similarity_search(
user_question,
k=vector_k,
)
# 3. 合并 + 去重
candidate_docs = []
seen_keys = set()
for doc in keyword_docs + semantic_docs:
unique_key = (
doc.metadata.get("source", ""),
doc.metadata.get("page", ""),
doc.page_content[:80],
)
if unique_key not in seen_keys:
candidate_docs.append(doc)
seen_keys.add(unique_key)
if not candidate_docs:
return []
# 4. LangChain 封装好的 rerank
reranked_docs = reranker.compress_documents(
documents=candidate_docs,
query=user_question,
)
return list(reranked_docs)[:final_k]
# =========================
# 7. 构造上下文
# =========================
def build_context(docs):
"""
把检索到的资料整理成 prompt 里的参考文档。
"""
context_list = []
source_list = []
for i, doc in enumerate(docs, start=1):
metadata = doc.metadata
source = metadata.get("source", "未知来源")
page = metadata.get("page", None)
if page is not None:
page_show = page + 1
else:
page_show = "未知页码"
context_list.append(
f"""
[资料{i}]
来源:{source}
页码:第 {page_show} 页
内容:
{doc.page_content}
"""
)
source_list.append({
"source": source,
"page": page_show,
})
context = "\n\n".join(context_list)
return context, source_list
# =========================
# 8. 问答
# =========================
def answer_question(
user_question,
llm,
bm25_retriever,
vector_store,
reranker,
):
"""
完整问答流程:
1. 混合检索
2. 拼接上下文
3. 调用大模型
4. 返回答案和来源
"""
docs = hybrid_search(
user_question=user_question,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
bm25_k=BM25_K,
vector_k=VECTOR_K,
final_k=FINAL_K,
reranker=reranker,
)
context, source_list = build_context(docs)
prompt = ChatPromptTemplate.from_template(
"""
你是一个专业的健康顾问,但不是医生,你只能根据参考文档回答。
要求:
1. 只根据参考文档回答。
2. 如果参考文档里没有答案,直接说:文档资料没有提到,我不知道。
3. 不要编造参考文档里没有的内容。
4. 你只是提供健康建议,不能代替医生诊断。
5. 回答要通俗易懂,说人话。
6. 不要在正文里标注来源。
参考文档:
{context}
用户问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
answer = chain.invoke({
"context": context,
"user_question": user_question,
})
source_text = "\n".join([
f"{i + 1}. {item.get('source', '未知来源')},第 {item.get('page', '未知页码')} 页"
for i, item in enumerate(source_list)
])
return f"{answer}\n\n参考文档:\n{source_text}"
# =========================
# 9. 菜单功能
# =========================
def run_ingest(vector_store):
"""
入库模式。
"""
raw_docs = load_pdf(PDF_PATH)
chunks = split_documents(
raw_docs,
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
)
ingest_documents(vector_store, chunks)
def run_qa(llm, vector_store, reranker):
"""
提问模式。
"""
bm25_retriever = load_bm25_retriever(k=BM25_K)
while True:
user_question = input("你想了解什么(输入 q 退出):").strip()
if user_question.lower() == "q":
print("已退出提问模式")
break
if not user_question:
print("问题不能为空")
continue
answer = answer_question(
user_question=user_question,
llm=llm,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
reranker=reranker,
)
print("\n" + answer + "\n")
def run_delete(vector_store):
"""
删除模式。
"""
confirm = input(f"确认删除 {COLLECTION_NAME} 吗?输入 yes:").strip()
if confirm == "yes":
delete_collection(vector_store)
else:
print("已取消删除")
# =========================
# 10. 主函数
# =========================
def main():
"""
主函数只做三件事:
1. 初始化环境
2. 初始化模型和数据库
3. 根据用户选择调用不同功能
"""
load_dotenv()
embedding = create_embedding_model()
vector_store = create_vector_store(embedding)
llm = create_llm()
reranker = create_reranker()
while True:
print("\n请选择模式:")
print("1 = 入库")
print("2 = 提问")
print("3 = 删除 collection")
print("q = 退出")
mode = input("请输入:").strip()
if mode == "1":
run_ingest(vector_store)
elif mode == "2":
run_qa(llm, vector_store, reranker)
elif mode == "3":
run_delete(vector_store)
elif mode.lower() == "q":
print("程序已退出")
break
else:
print("请输入 1、2、3 或 q")
if __name__ == "__main__":
main()
返回结果:
你想了解什么(输入 q 退出):青少年生长缓慢怎么办
青少年生长缓慢,首先要从饮食、运动、睡眠和心理等方面进行综合调理。
1. **吃好每一顿饭**:保证一日三餐规律,食物种类要丰富。每天至少吃12种以上的食物,每周达到25种以上更好。主食不要只吃白米白面,可以搭配杂粮、薯类;肉类可以轮流吃瘦肉、禽肉、鱼虾;多吃奶制品、蛋类和豆制品,这些是优质蛋白和钙的好来源。新鲜蔬菜水果也要足量吃。
2. **适当增加营养**:如果存在挑食、偏食的情况,可以在平衡饮食的基础上,适当多吃一些富含优质蛋白质的食物,比如瘦肉、鱼、蛋、大豆等。每天喝300毫升以上的牛奶或吃相应的奶制品,同时注意补充富含维生素D的食物,必要时可在专业人士指导下补充维生素D。
3. **合理运动**:每天保持适量的身体活动,减少久坐时间,限制看电子屏幕的时间(6~17岁不超过2小时,越少越好)。规律运动有助于促进食欲和骨骼发育。
4. **保证充足睡眠**:睡眠对长高很重要。13~17岁的青少年每天应睡8~10小时。睡得好,生长激素才能正常分泌。
5. **关注心理和情绪**:避免因情绪问题导致的少吃或限制进食。家长要帮助孩子正确认识体型,保持健康体重,有需要时给予心理支持。
6. **定期监测身高体重**:至少每半年观察一次体格发育情况,发现问题及时调整饮食或寻求专业指导。
如果长期生长发育不理想,改善效果不明显,或者怀疑有疾病原因,建议及时去医院检查,排除病理性因素。
这些建议适用于因营养不足引起的生长缓慢,不能代替医生诊断。具体情况最好在医师或营养指导人员的帮助下制定个性化方案。
参考文档:
1. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 12 页
2. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 6 页
3. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 11 页
4. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 7 页
5. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 4 页
(4)query rewrite 查询改写
查询改写,很容易理解,把用户随口说的、模糊的、太短的问题,改写成更适合去资料库搜索的问题。就比如我这里,如果用户询问"我的孩子怎么长不高啊?"->"青少年生长缓慢是什么因素导致的,并提出相应解决方案",也就是把不好检索的问题,变得更有针对性。
值得注意的是有时候我们的改写需要用到上下文,所以这里先介绍一种langChain封装的,记忆方法,这里仅仅介绍暂时记忆,主要为了记录上下文5轮,其实很简单 核心只需要一个memory记忆器即可实现
#创建memory
from langchain_classic.memory import ConversationBufferWindowMemory
memory = ConversationBufferWindowMemory(
k=5, #保存多少轮
return_messages=False, #用于返回文本 不要返回message类型
memory_key="chat_history",
input_key="user_question",
)
'''
message类型 messages[-1].content = "memory 可以理解成聊天记忆本"
[
HumanMessage(content="你好"),
AIMessage(content="你好呀,有什么我可以帮你?"),
HumanMessage(content="帮我解释一下 memory"),
AIMessage(content="memory 可以理解成聊天记忆本")
]
文本类型
Human: 你好
AI: 你好呀
Human: 帮我解释 memory
AI: memory 是聊天记忆本
'''
点击查看完整版代码和返回结果
import os
import pickle
import jieba
from dotenv import load_dotenv
from langchain_community.document_loaders import PyPDFLoader
from langchain_community.retrievers import BM25Retriever
from langchain_community.document_compressors.dashscope_rerank import DashScopeRerank
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_chroma import Chroma
from langchain_openai import OpenAIEmbeddings, ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_classic.memory import ConversationBufferWindowMemory
# =========================
# 1. 全局配置
# =========================
PDF_PATH = r"C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf"
BM25_DOCS_PATH = "./bm25_docs.pkl"
COLLECTION_NAME = "my_rag_collection"
DB_PATH = "./chroma_db"
CHUNK_SIZE = 500
CHUNK_OVERLAP = 100
BM25_K = 10
VECTOR_K = 10
FINAL_K = 5
# =========================
# 2. 文档加载与切分
# =========================
def load_pdf(path):
"""
读取 PDF。
返回 LangChain 的 Document 列表。
"""
loader = PyPDFLoader(path)
return loader.load()
def split_documents(raw_docs, chunk_size=500, chunk_overlap=100):
"""
把大文档切成小块。
chunk 说人话:就是一小段资料。
"""
splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
separators=[
"\n\n",
"\n",
"。",
"!",
"?",
";",
",",
"、",
" ",
"",
],
is_separator_regex=False,
)
chunks = splitter.split_documents(raw_docs)
print(f"切分出 {len(chunks)} 个 chunk")
return chunks
# =========================
# 3. 初始化模型和数据库
# =========================
def create_embedding_model():
"""
创建 embedding 模型。
embedding 说人话:把文字变成数字,方便机器比较相似度。
"""
return OpenAIEmbeddings(
model=os.getenv("DASHSCOPE_EMBEDDING_MODEL"),
api_key=os.getenv("DASHSCOPE_API_KEY"),
base_url=os.getenv("DASHSCOPE_BASE_URL"),
dimensions=1024,
chunk_size=10,
check_embedding_ctx_length=False,
)
def create_vector_store(embedding):
"""
创建 Chroma 向量数据库。
Chroma 说人话:本地知识库。
"""
return Chroma(
collection_name=COLLECTION_NAME,
embedding_function=embedding,
persist_directory=DB_PATH,
)
def create_llm():
"""
创建大模型。
"""
return ChatOpenAI(
model=os.getenv("DASHSCOPE_CHAT_MODEL"),
api_key=os.getenv("DASHSCOPE_API_KEY"),
base_url=os.getenv("DASHSCOPE_BASE_URL"),
)
def create_reranker():
"""
创建 rerank 重排器。
rerank 说人话:把已经找出来的资料重新排队,把最有用的放前面。
"""
return DashScopeRerank(
model=os.getenv("DASHSCOPE_RERANK_MODEL"),
top_n=FINAL_K,
dashscope_api_key=os.getenv("DASHSCOPE_API_KEY"),
)
def create_memory():
"""
创建上下文短期记忆。
只记最近 5 轮对话。
"""
return ConversationBufferWindowMemory(
k=5,
return_messages=False,
memory_key="chat_history",
input_key="user_question",
)
# =========================
# 4. BM25 相关
# =========================
def chinese_tokenizer(text):
"""
给 BM25 用的中文分词函数。
分词说人话:把一句中文切成一个个词。
"""
return [word.strip() for word in jieba.cut(text) if word.strip()]
def save_bm25_docs(chunks):
"""
保存 chunks,给 BM25 使用。
"""
with open(BM25_DOCS_PATH, "wb") as f:
pickle.dump(chunks, f)
print(f"BM25 文档已保存到 {BM25_DOCS_PATH}")
def load_bm25_retriever(k=5):
"""
加载 BM25 检索器。
BM25 说人话:按关键词找资料。
"""
if not os.path.exists(BM25_DOCS_PATH):
raise FileNotFoundError(
f"找不到 {BM25_DOCS_PATH},请先选择入库模式。"
)
with open(BM25_DOCS_PATH, "rb") as f:
docs = pickle.load(f)
retriever = BM25Retriever.from_documents(
docs,
preprocess_func=chinese_tokenizer,
)
retriever.k = k
return retriever
# =========================
# 5. 入库、删除
# =========================
def ingest_documents(vector_store, chunks):
"""
入库:
1. 存进 Chroma,给向量检索用。
2. 存成 pkl,给 BM25 用。
"""
ids = [f"text_new_chunk_{i}" for i in range(len(chunks))]
vector_store.add_documents(
documents=chunks,
ids=ids,
)
save_bm25_docs(chunks)
print(f"成功入库 {len(chunks)} 个 chunk")
def delete_collection(vector_store):
"""
删除 Chroma collection 和 BM25 文件。
collection 说人话:知识库里的一张表。
"""
vector_store.delete_collection()
print(f"{COLLECTION_NAME} 已删除")
if os.path.exists(BM25_DOCS_PATH):
os.remove(BM25_DOCS_PATH)
print(f"{BM25_DOCS_PATH} 已删除")
# =========================
# 6. 混合检索
# =========================
def hybrid_search(
user_question,
bm25_retriever,
vector_store,
reranker,
bm25_k=10,
vector_k=10,
final_k=5,
):
"""
混合检索 + LangChain rerank:
1. BM25 按关键词找
2. Chroma 按意思找
3. 合并
4. 去重
5. 用 LangChain 封装好的 DashScopeRerank 重排
"""
bm25_retriever.k = bm25_k
# 1. 关键词检索
keyword_docs = bm25_retriever.invoke(user_question)
# 2. 语义检索
semantic_docs = vector_store.similarity_search(
user_question,
k=vector_k,
)
# 3. 合并 + 去重
candidate_docs = []
seen_keys = set()
for doc in keyword_docs + semantic_docs:
unique_key = (
doc.metadata.get("source", ""),
doc.metadata.get("page", ""),
doc.page_content[:80],
)
if unique_key not in seen_keys:
candidate_docs.append(doc)
seen_keys.add(unique_key)
if not candidate_docs:
return []
# 4. LangChain 封装好的 rerank
reranked_docs = reranker.compress_documents(
documents=candidate_docs,
query=user_question,
)
return list(reranked_docs)[:final_k]
# =========================
# 7. 构造上下文
# =========================
def build_context(docs):
"""
把检索到的资料整理成 prompt 里的参考文档。
"""
context_list = []
source_list = []
for i, doc in enumerate(docs, start=1):
metadata = doc.metadata
source = metadata.get("source", "未知来源")
page = metadata.get("page", None)
if page is not None:
page_show = page + 1
else:
page_show = "未知页码"
context_list.append(
f"""
[资料{i}]
来源:{source}
页码:第 {page_show} 页
内容:
{doc.page_content}
"""
)
source_list.append({
"source": source,
"page": page_show,
})
context = "\n\n".join(context_list)
return context, source_list
# =========================
# 8. 问答
# =========================
def answer_question(
user_question,
llm,
bm25_retriever,
vector_store,
reranker,
memory,
):
"""
完整问答流程:
1. 混合检索
2. 拼接上下文
3. 读取最近 5 轮历史对话
4. 调用大模型
5. 保存本轮问答到记忆
6. 返回答案和来源
"""
docs = hybrid_search(
user_question=user_question,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
bm25_k=BM25_K,
vector_k=VECTOR_K,
final_k=FINAL_K,
reranker=reranker,
)
context, source_list = build_context(docs)
chat_history = memory.load_memory_variables({}).get("chat_history", "")
prompt = ChatPromptTemplate.from_template(
"""
你是一个专业的健康顾问,但不是医生,你只能根据参考文档回答。
要求:
1. 主要根据参考文档回答。
2. 历史对话只用于理解用户当前问题,不要当成资料来源。
3. 如果参考文档里没有答案,直接说:文档资料没有提到,我不知道。
4. 不要编造参考文档里没有的内容。
5. 你只是提供健康建议,不能代替医生诊断。
6. 回答要通俗易懂,说人话。
7. 不要在正文里标注来源。
历史对话:
{chat_history}
参考文档:
{context}
用户问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
answer = chain.invoke({
"chat_history": chat_history,
"context": context,
"user_question": user_question,
})
memory.save_context(
{"user_question": user_question},
{"output": answer},
)
source_text = "\n".join([
f"{i + 1}. {item.get('source', '未知来源')},第 {item.get('page', '未知页码')} 页"
for i, item in enumerate(source_list)
])
return f"{answer}\n\n参考文档:\n{source_text}"
# =========================
# 9. 菜单功能
# =========================
def run_ingest(vector_store):
"""
入库模式。
"""
raw_docs = load_pdf(PDF_PATH)
chunks = split_documents(
raw_docs,
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
)
ingest_documents(vector_store, chunks)
def run_qa(llm, vector_store, reranker, memory):
"""
提问模式。
"""
bm25_retriever = load_bm25_retriever(k=BM25_K)
while True:
user_question = input("你想了解什么(输入 q 退出,输入 clear 清空记忆):").strip()
if user_question.lower() == "q":
print("已退出提问模式")
break
if user_question.lower() == "clear":
memory.clear()
print("已清空上下文记忆")
continue
if not user_question:
print("问题不能为空")
continue
answer = answer_question(
user_question=user_question,
llm=llm,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
reranker=reranker,
memory=memory,
)
print("\n" + answer + "\n")
def run_delete(vector_store):
"""
删除模式。
"""
confirm = input(f"确认删除 {COLLECTION_NAME} 吗?输入 yes:").strip()
if confirm == "yes":
delete_collection(vector_store)
else:
print("已取消删除")
# =========================
# 10. 主函数
# =========================
def main():
"""
主函数只做三件事:
1. 初始化环境
2. 初始化模型和数据库
3. 根据用户选择调用不同功能
"""
load_dotenv()
embedding = create_embedding_model()
vector_store = create_vector_store(embedding)
llm = create_llm()
reranker = create_reranker()
memory = create_memory()
while True:
print("\n请选择模式:")
print("1 = 入库")
print("2 = 提问")
print("3 = 删除 collection")
print("q = 退出")
mode = input("请输入:").strip()
if mode == "1":
run_ingest(vector_store)
elif mode == "2":
run_qa(llm, vector_store, reranker, memory)
elif mode == "3":
run_delete(vector_store)
elif mode.lower() == "q":
print("程序已退出")
break
else:
print("请输入 1、2、3 或 q")
if __name__ == "__main__":
main()
返回结果:
你想了解什么(输入 q 退出):我的孩子长不高是什么原因
孩子长不高可能和长期营养摄入不均衡有关,比如蛋白质、维生素或矿物质吃得不够。如果身高明显低于同龄人,可能是生长迟缓,这通常是因为平时吃饭不规律、挑食、偏食,或者饮食结构不合理导致的。
另外,睡眠不足、运动少、情绪不好也可能影响孩子的生长发育。中医认为,脾胃功能弱、食欲差、消化吸收不好,也会影响长个子。
建议先看看孩子平时吃饭、睡觉、运动情况怎么样,有没有挑食或睡不好等问题。如果有疑虑,最好找专业医生或营养师评估一下,不要自行判断。
记住,这不是医疗诊断,只能作为日常参考。
参考文档:
1. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 4 页
2. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 11 页
3. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 4 页
4. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 3 页
5. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 7 页
你想了解什么(输入 q 退出):怎么办呢
如果孩子长不高,首先要从日常饮食、作息和生活习惯入手来改善。
1. **保证营养均衡**:每天吃够多种食物,包括主食、蔬菜水果、肉蛋奶豆类等。特别要注意钙、维生素D、铁、锌、碘这些对长个子很重要的营养素。可以多吃奶制品、深色蔬菜、动物肝脏、瘦肉、鱼虾等。
2. **三餐规律**:按时吃饭,不挑食、不偏食,避免暴饮暴食或节食,保护好脾胃功能。
3. **多运动**:每天适当进行跳绳、打球、跑步、游泳等有助于长高的运动,促进骨骼发育。
4. **睡得好**:保证充足的睡眠,尽量在晚上10点前入睡,因为生长激素在深度睡眠时分泌最多。
5. **心情愉快**:压力大、情绪差也会影响长高,要关注孩子的心理健康。
6. **定期检查**:建议定期测量身高体重,看看是否在正常范围。如果长期比同龄人矮很多,或者改善一段时间后效果不明显,应该去医院找专业医生评估,看有没有疾病或其他原因。
记住,这些建议只是日常调理参考,不能代替医生的诊断和治疗。
参考文档:
1. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 12 页
2. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 6 页
3. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 4 页
4. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 7 页
5. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 3 页
记忆器已经实现完成,下一步就是进行query rewrite查询改写
#核心代码
#构造rewrite函数
def rewrite_query(user_question,llm,memory):
#加载记忆
chat_history = memory.load_memory_variables({}).get("chat_history")
prompt = ChatPromptTemplate.from_template(
"""
你需要把用户当前问题改写成一个独立、完整、适合检索资料的问题。
要求:
1. 只输出改写后的问题,不要解释。
2. 不要回答问题。
3. 不要编造历史对话和用户问题里没有的信息。
4. 如果当前问题本身已经完整,就原样输出。
历史对话:
{chat_history}
用户当前问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
rewritten_question = chain.invoke({
"chat_history": chat_history,
"user_question": user_question,
})
#这个模型可以更换 这里为方便未换
rewritten_question = rewritten_question.strip()
print(f"原问题:{user_question}")
print(f"改写后问题:{rewritten_question}")
#方便检查修改
return rewritten_question
点击查看完整代码和返回结果
import os
import pickle
import jieba
from dotenv import load_dotenv
from langchain_community.document_loaders import PyPDFLoader
from langchain_community.retrievers import BM25Retriever
from langchain_community.document_compressors.dashscope_rerank import DashScopeRerank
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_chroma import Chroma
from langchain_openai import OpenAIEmbeddings, ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_classic.memory import ConversationBufferWindowMemory
from pydantic.v1 import SecretStr
# =========================
# 1. 全局配置
# =========================
PDF_PATH = r"C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf"
BM25_DOCS_PATH = "./bm25_docs.pkl"
COLLECTION_NAME = "my_rag_collection"
DB_PATH = "./chroma_db"
CHUNK_SIZE = 500
CHUNK_OVERLAP = 100
BM25_K = 10
VECTOR_K = 10
FINAL_K = 5
# =========================
# 2. 文档加载与切分
# =========================
def load_pdf(path):
"""
读取 PDF。
返回 LangChain 的 Document 列表。
"""
loader = PyPDFLoader(path)
return loader.load()
def split_documents(raw_docs, chunk_size=500, chunk_overlap=100):
"""
把大文档切成小块。
chunk 说人话:就是一小段资料。
"""
splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
separators=[
"\n\n",
"\n",
"。",
"!",
"?",
";",
",",
"、",
" ",
"",
],
is_separator_regex=False,
)
chunks = splitter.split_documents(raw_docs)
print(f"切分出 {len(chunks)} 个 chunk")
return chunks
# =========================
# 3. 初始化模型和数据库
# =========================
def create_embedding_model():
"""
创建 embedding 模型。
embedding 说人话:把文字变成数字,方便机器比较相似度。
"""
return OpenAIEmbeddings(
model=os.getenv("DASHSCOPE_EMBEDDING_MODEL"),
api_key=os.getenv("DASHSCOPE_API_KEY", ""),
base_url=os.getenv("DASHSCOPE_BASE_URL"),
dimensions=1024,
chunk_size=10,
check_embedding_ctx_length=False,
)
def create_vector_store(embedding):
"""
创建 Chroma 向量数据库。
Chroma 说人话:本地知识库。
"""
return Chroma(
collection_name=COLLECTION_NAME,
embedding_function=embedding,
persist_directory=DB_PATH,
)
def create_llm():
"""
创建大模型。
"""
return ChatOpenAI(
model=os.getenv("DASHSCOPE_CHAT_MODEL"),
api_key=os.getenv("DASHSCOPE_API_KEY"),
base_url=os.getenv("DASHSCOPE_BASE_URL"),
)
def create_reranker():
"""
创建 rerank 重排器。
rerank 说人话:把已经找出来的资料重新排队,把最有用的放前面。
"""
return DashScopeRerank(
model=os.getenv("DASHSCOPE_RERANK_MODEL"),
top_n=FINAL_K,
)
def create_memory():
"""
创建上下文短期记忆。
只记最近 5 轮对话。
"""
return ConversationBufferWindowMemory(
k=5,
return_messages=False,
memory_key="chat_history",
input_key="user_question",
)
def rewrite_query(user_question, llm, memory):
"""
重写器,拿到用户的问题进行重写
"""
chat_history = memory.load_memory_variables({}).get("chat_history")
prompt = ChatPromptTemplate.from_template(
"""
你需要把用户当前问题改写成一个独立、完整、适合检索资料的问题。
要求:
1. 只输出改写后的问题,不要解释。
2. 不要回答问题。
3. 不要编造历史对话和用户问题里没有的信息。
4. 如果当前问题本身已经完整,就原样输出。
历史对话:
{chat_history}
用户当前问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
rewritten_question = chain.invoke({
"chat_history": chat_history,
"user_question": user_question,
})
rewritten_question = rewritten_question.strip()
print(f"原问题:{user_question}")
print(f"改写后问题:{rewritten_question}")
return rewritten_question
# =========================
# 4. BM25 相关
# =========================
def chinese_tokenizer(text):
"""
给 BM25 用的中文分词函数。
分词说人话:把一句中文切成一个个词。
"""
return [word.strip() for word in jieba.cut(text) if word.strip()]
def save_bm25_docs(chunks):
"""
保存 chunks,给 BM25 使用。
"""
with open(BM25_DOCS_PATH, "wb") as f:
pickle.dump(chunks, f)
print(f"BM25 文档已保存到 {BM25_DOCS_PATH}")
def load_bm25_retriever(k=5):
"""
加载 BM25 检索器。
BM25 说人话:按关键词找资料。
"""
if not os.path.exists(BM25_DOCS_PATH):
raise FileNotFoundError(
f"找不到 {BM25_DOCS_PATH},请先选择入库模式。"
)
with open(BM25_DOCS_PATH, "rb") as f:
docs = pickle.load(f)
retriever = BM25Retriever.from_documents(
docs,
preprocess_func=chinese_tokenizer,
)
retriever.k = k
return retriever
# =========================
# 5. 入库、删除
# =========================
def ingest_documents(vector_store, chunks):
"""
入库:
1. 存进 Chroma,给向量检索用。
2. 存成 pkl,给 BM25 用。
"""
ids = [f"text_new_chunk_{i}" for i in range(len(chunks))]
vector_store.add_documents(
documents=chunks,
ids=ids,
)
save_bm25_docs(chunks)
print(f"成功入库 {len(chunks)} 个 chunk")
def delete_collection(vector_store):
"""
删除 Chroma collection 和 BM25 文件。
collection 说人话:知识库里的一张表。
"""
vector_store.delete_collection()
print(f"{COLLECTION_NAME} 已删除")
if os.path.exists(BM25_DOCS_PATH):
os.remove(BM25_DOCS_PATH)
print(f"{BM25_DOCS_PATH} 已删除")
# =========================
# 6. 混合检索
# =========================
def hybrid_search(
user_question,
bm25_retriever,
vector_store,
reranker,
bm25_k=10,
vector_k=10,
final_k=5,
):
"""
混合检索 + LangChain rerank:
1. BM25 按关键词找
2. Chroma 按意思找
3. 合并
4. 去重
5. 用 LangChain 封装好的 DashScopeRerank 重排
"""
bm25_retriever.k = bm25_k
# 1. 关键词检索
keyword_docs = bm25_retriever.invoke(user_question)
# 2. 语义检索
semantic_docs = vector_store.similarity_search(
user_question,
k=vector_k,
)
# 3. 合并 + 去重
candidate_docs = []
seen_keys = set()
for doc in keyword_docs + semantic_docs:
unique_key = (
doc.metadata.get("source", ""),
doc.metadata.get("page", ""),
doc.page_content[:80],
)
if unique_key not in seen_keys:
candidate_docs.append(doc)
seen_keys.add(unique_key)
if not candidate_docs:
return []
# 4. LangChain 封装好的 rerank
reranked_docs = reranker.compress_documents(
documents=candidate_docs,
query=user_question,
)
return list(reranked_docs)[:final_k]
# =========================
# 7. 构造上下文
# =========================
def build_context(docs):
"""
把检索到的资料整理成 prompt 里的参考文档。
"""
context_list = []
source_list = []
for i, doc in enumerate(docs, start=1):
metadata = doc.metadata
source = metadata.get("source", "未知来源")
page = metadata.get("page", None)
if page is not None:
page_show = page + 1
else:
page_show = "未知页码"
context_list.append(
f"""
[资料{i}]
来源:{source}
页码:第 {page_show} 页
内容:
{doc.page_content}
"""
)
source_list.append({
"source": source,
"page": page_show,
})
context = "\n\n".join(context_list)
return context, source_list
# =========================
# 8. 问答
# =========================
def answer_question(
user_question,
llm,
bm25_retriever,
vector_store,
reranker,
memory,
):
"""
完整问答流程:
1. 查询重写
2. 混合检索
3. 拼接上下文
4. 读取最近 5 轮历史对话
5. 调用大模型
6. 保存本轮问答到记忆
7. 返回答案和来源
"""
rewritten_question = rewrite_query(
user_question=user_question,
llm=llm,
memory=memory,
)
docs = hybrid_search(
user_question=rewritten_question,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
bm25_k=BM25_K,
vector_k=VECTOR_K,
final_k=FINAL_K,
reranker=reranker,
)
context, source_list = build_context(docs)
chat_history = memory.load_memory_variables({}).get("chat_history", "")
prompt = ChatPromptTemplate.from_template(
"""
你是一个专业的健康顾问,但不是医生,你只能根据参考文档回答。
要求:
1. 主要根据参考文档回答。
2. 历史对话只用于理解用户当前问题,不要当成资料来源。
3. 如果参考文档里没有答案,直接说:文档资料没有提到,我不知道。
4. 不要编造参考文档里没有的内容。
5. 你只是提供健康建议,不能代替医生诊断。
6. 回答要通俗易懂,说人话。
7. 不要在正文里标注来源。
历史对话:
{chat_history}
参考文档:
{context}
用户问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
answer = chain.invoke({
"chat_history": chat_history,
"context": context,
"user_question": user_question,
})
memory.save_context(
{"user_question": user_question},
{"output": answer},
)
source_text = "\n".join([
f"{i + 1}. {item.get('source', '未知来源')},第 {item.get('page', '未知页码')} 页"
for i, item in enumerate(source_list)
])
return f"{answer}\n\n参考文档:\n{source_text}"
# =========================
# 9. 菜单功能
# =========================
def run_ingest(vector_store):
"""
入库模式。
"""
raw_docs = load_pdf(PDF_PATH)
chunks = split_documents(
raw_docs,
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
)
ingest_documents(vector_store, chunks)
def run_qa(llm, vector_store, reranker, memory):
"""
提问模式。
"""
bm25_retriever = load_bm25_retriever(k=BM25_K)
while True:
user_question = input("你想了解什么(输入 q 退出,输入 clear 清空记忆):").strip()
if user_question.lower() == "q":
print("已退出提问模式")
break
if user_question.lower() == "clear":
memory.clear()
print("已清空上下文记忆")
continue
if not user_question:
print("问题不能为空")
continue
answer = answer_question(
user_question=user_question,
llm=llm,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
reranker=reranker,
memory=memory,
)
print("\n" + answer + "\n")
def run_delete(vector_store):
"""
删除模式。
"""
confirm = input(f"确认删除 {COLLECTION_NAME} 吗?输入 yes:").strip()
if confirm == "yes":
delete_collection(vector_store)
else:
print("已取消删除")
# =========================
# 10. 主函数
# =========================
def main():
"""
主函数只做三件事:
1. 初始化环境
2. 初始化模型和数据库
3. 根据用户选择调用不同功能
"""
load_dotenv()
embedding = create_embedding_model()
vector_store = create_vector_store(embedding)
llm = create_llm()
reranker = create_reranker()
memory = create_memory()
while True:
print("\n请选择模式:")
print("1 = 入库")
print("2 = 提问")
print("3 = 删除 collection")
print("q = 退出")
mode = input("请输入:").strip()
if mode == "1":
run_ingest(vector_store)
elif mode == "2":
run_qa(llm, vector_store, reranker, memory)
elif mode == "3":
run_delete(vector_store)
elif mode.lower() == "q":
print("程序已退出")
break
else:
print("请输入 1、2、3 或 q")
if __name__ == "__main__":
main()
返回结果:
你想了解什么(输入 q 退出,输入 clear 清空记忆):青少年生长缓慢怎么办
原问题:青少年生长缓慢怎么办
改写后问题:青少年生长缓慢怎么办
青少年生长缓慢,首先要从饮食、运动、睡眠和心理等多方面入手,帮助改善生长发育状况。
1. **保证营养全面**:一日三餐要规律,食物种类要多样。每天至少吃12种以上不同的食物,每周尽量达到25种以上。主食不要只吃白米白面,可以搭配杂粮、薯类;肉类可轮换着吃畜肉、禽肉和鱼虾;蛋类、奶制品和大豆也要经常吃。特别是要多吃富含优质蛋白质的食物,比如瘦肉、鱼、蛋、奶和豆制品。
2. **保证奶制品摄入**:每天喝300毫升以上的牛奶或吃等量的酸奶、奶酪等,同时注意补充富含维生素D的食物(如蛋黄、深海鱼),必要时可在专业人士指导下补充维生素D。
3. **多吃新鲜蔬果**:蔬菜水果要足量,尤其是深色蔬菜和应季水果,有助于补充维生素和矿物质。
4. **避免挑食偏食**:如果孩子有挑食、偏食的情况,家长要注意引导,通过同类食物互换(比如鸡肉换鸭肉、鸡蛋换鹌鹑蛋)来增加食物多样性。
5. **合理加餐**:如果孩子食欲差或饭量小,可以在正餐之间适当加餐,选择营养密度高的食物,比如奶类、水果、坚果等。
6. **加强身体活动**:每天进行适量运动,比如跳绳、打球、跑步、游泳等,有助于促进骨骼生长和食欲。建议限制看电子屏幕的时间,6~17岁孩子每天不超过2小时,越少越好。
7. **保证充足睡眠**:睡眠对长高很重要。13~17岁的青少年每天应睡够8~10小时。睡得好,生长激素才能正常分泌。
8. **关注心理健康**:情绪不好、压力大也可能影响吃饭和发育。要关心孩子的心理状态,避免因体型问题产生焦虑,造成节食或进食异常。
9. **定期监测身高体重**:至少每半年测量一次身高体重,观察生长趋势。如果长期长得慢、改善不明显,应及时去医院检查,排除疾病因素。
需要注意的是,如果是由于慢性病、内分泌问题(如生长激素缺乏)等原因引起的生长迟缓,不能单靠饮食调理,必须在医生指导下治疗。
总之,先从日常饮食和生活习惯做起,必要时寻求专业医师或营养指导人员的帮助,制定个性化方案。我不是医生,建议仅供参考,不能代替医疗诊断。
参考文档:
1. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 12 页
2. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 6 页
3. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 11 页
4. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 7 页
5. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 4 页
你想了解什么(输入 q 退出,输入 clear 清空记忆):他还有高血压
原问题:他还有高血压
改写后问题:青少年生长缓慢且伴有高血压应如何处理?
文档资料没有提到,我不知道。
参考文档:
1. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 12 页
2. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 6 页
3. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 11 页
4. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 7 页
5. C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf,第 4 页
(5)RAG评估
到这里主要的常见的通用优化基本已经讲解完毕,但是对于自己的项目某些细节上还需要继续优化,这里附上完整代码,从代码中去对于各个部分进行更精细的优化,这里的优化点可以通过ai实现,因为大多是小细节问题,一点点抠代码细节对于个人而言成本太高。
点击查看代码
import os
import re
import json
import pickle
import hashlib
from pathlib import Path
from typing import List, Dict, Tuple
import jieba
from dotenv import load_dotenv
from langchain_community.document_loaders import PyPDFLoader
from langchain_community.retrievers import BM25Retriever
from langchain_community.document_compressors.dashscope_rerank import DashScopeRerank
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_chroma import Chroma
from langchain_openai import OpenAIEmbeddings, ChatOpenAI
from langchain_core.documents import Document
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
# =========================
# 0. 加载环境变量
# =========================
# ✅ 优化点:提前 load_dotenv,避免下面读取 os.getenv 时读不到 .env
load_dotenv()
# =========================
# 1. 全局配置
# =========================
PDF_PATH = os.getenv(
"PDF_PATH",
r"C:\Users\向梒宇\Desktop\health_konwlege\通用\儿童青少年生长迟缓食养指南 2023版.pdf"
)
BM25_DOCS_PATH = os.getenv("BM25_DOCS_PATH", "./bm25_docs.pkl")
COLLECTION_NAME = os.getenv("COLLECTION_NAME", "my_rag_collection")
DB_PATH = os.getenv("DB_PATH", "./chroma_db")
CHUNK_SIZE = int(os.getenv("CHUNK_SIZE", "500"))
CHUNK_OVERLAP = int(os.getenv("CHUNK_OVERLAP", "100"))
BM25_K = int(os.getenv("BM25_K", "10"))
VECTOR_K = int(os.getenv("VECTOR_K", "10"))
FINAL_K = int(os.getenv("FINAL_K", "5"))
MEMORY_K = int(os.getenv("MEMORY_K", "5"))
# =========================
# 2. 工具函数
# =========================
def require_env(name: str) -> str:
"""
读取必须存在的环境变量。
如果没配,直接报错。
"""
value = os.getenv(name)
if not value:
raise RuntimeError(f"缺少环境变量:{name},请检查 .env 文件")
return value
def stable_doc_id(doc: Document, index: int) -> str:
"""
生成稳定 ID。
hash 说人话:把一段内容压成一个短指纹,用来区分它是谁。
"""
source = doc.metadata.get("source", "")
page = str(doc.metadata.get("page", ""))
text = doc.page_content
raw = f"{source}|{page}|{index}|{text}"
digest = hashlib.md5(raw.encode("utf-8")).hexdigest()
return f"chunk_{digest}"
def clean_text(text: str) -> str:
"""
清洗 PDF 文字。
作用:减少多余空格、奇怪换行。
"""
if not text:
return ""
text = text.replace("\u3000", " ")
text = re.sub(r"[ \t]+", " ", text)
text = re.sub(r"\n{3,}", "\n\n", text)
text = text.strip()
return text
def is_low_quality_text(text: str) -> bool:
"""
判断是不是低质量片段。
低质量片段说人话:目录、页码、太短的废内容。
"""
text = text.strip()
if len(text) < 60:
return True
# ✅ 优化点:过滤目录类 chunk
# 例如:一堆 .......... 6、............. 12
dot_count = text.count(".") + text.count("·") + text.count("…")
if dot_count >= 10 and ("目录" in text or "附录" in text):
return True
# ✅ 优化点:过滤明显目录页
if "目录" in text and re.search(r"\.{3,}|…{2,}", text):
return True
return False
def page_show(doc: Document):
"""
LangChain 里的 page 通常从 0 开始。
展示给用户时 +1。
"""
page = doc.metadata.get("page", None)
if page is None:
return "未知页码"
return page + 1
# =========================
# 3. 简单窗口记忆
# =========================
class SimpleWindowMemory:
"""
✅ 优化点:替代 ConversationBufferWindowMemory,避免 LangChain 废弃警告。
记忆说人话:只记最近几轮聊天,太早的就忘掉。
"""
def __init__(self, k: int = 5):
self.k = k
self.history: List[Dict[str, str]] = []
def load(self) -> str:
recent = self.history[-self.k:]
lines = []
for item in recent:
lines.append(f"用户:{item['user']}")
lines.append(f"助手:{item['assistant']}")
return "\n".join(lines)
def save(self, user_question: str, answer: str):
self.history.append({
"user": user_question,
"assistant": answer,
})
# 只保留最近 k 轮
self.history = self.history[-self.k:]
def clear(self):
self.history.clear()
# =========================
# 4. 文档加载与切分
# =========================
def load_pdf(path: str) -> List[Document]:
"""
读取 PDF。
"""
pdf_file = Path(path)
# ✅ 优化点:先检查文件是否存在,错误更清楚
if not pdf_file.exists():
raise FileNotFoundError(f"PDF 文件不存在:{path}")
loader = PyPDFLoader(str(pdf_file))
raw_docs = loader.load()
cleaned_docs = []
for doc in raw_docs:
text = clean_text(doc.page_content)
if not text:
continue
doc.page_content = text
cleaned_docs.append(doc)
print(f"读取 PDF 页数:{len(cleaned_docs)}")
return cleaned_docs
def split_documents(
raw_docs: List[Document],
chunk_size: int = 500,
chunk_overlap: int = 100,
) -> List[Document]:
"""
把大文档切成小块。
chunk 说人话:就是一小段资料。
"""
splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
separators=[
"\n\n",
"\n",
"。",
"!",
"?",
";",
",",
"、",
" ",
"",
],
is_separator_regex=False,
)
chunks = splitter.split_documents(raw_docs)
# ✅ 优化点:切分后再次清洗和过滤,减少目录、废片段进入知识库
good_chunks = []
for doc in chunks:
doc.page_content = clean_text(doc.page_content)
if is_low_quality_text(doc.page_content):
continue
good_chunks.append(doc)
print(f"原始切分 chunk 数:{len(chunks)}")
print(f"过滤后有效 chunk 数:{len(good_chunks)}")
return good_chunks
# =========================
# 5. 初始化模型和数据库
# =========================
def create_embedding_model():
"""
创建 embedding 模型。
embedding 说人话:把文字变成数字,方便机器比较相似度。
"""
return OpenAIEmbeddings(
model=require_env("DASHSCOPE_EMBEDDING_MODEL"),
api_key=require_env("DASHSCOPE_API_KEY"),
base_url=require_env("DASHSCOPE_BASE_URL"),
# ✅ 优化点:维度从环境变量读取,后续换模型更方便
dimensions=int(os.getenv("EMBEDDING_DIMENSIONS", "1024")),
chunk_size=10,
check_embedding_ctx_length=False,
)
def create_vector_store(embedding):
"""
创建 Chroma 向量数据库。
Chroma 说人话:本地知识库。
"""
return Chroma(
collection_name=COLLECTION_NAME,
embedding_function=embedding,
persist_directory=DB_PATH,
)
def create_llm():
"""
创建大模型。
"""
return ChatOpenAI(
model=require_env("DASHSCOPE_CHAT_MODEL"),
api_key=require_env("DASHSCOPE_API_KEY"),
base_url=require_env("DASHSCOPE_BASE_URL"),
# ✅ 优化点:RAG 问答建议温度设低,减少胡编
temperature=0,
# ✅ 优化点:增加超时和重试,网络抖动时更稳
timeout=60,
max_retries=2,
)
def create_reranker():
"""
创建 rerank 重排器。
rerank 说人话:把已经找出来的资料重新排队,把最有用的放前面。
"""
return DashScopeRerank(
model=require_env("DASHSCOPE_RERANK_MODEL"),
top_n=FINAL_K,
)
def create_memory():
"""
创建短期记忆。
"""
return SimpleWindowMemory(k=MEMORY_K)
# =========================
# 6. 查询改写与检索词扩展
# =========================
def rewrite_query(user_question: str, llm, memory: SimpleWindowMemory) -> str:
"""
把用户问题改写成适合检索的问题。
"""
chat_history = memory.load()
prompt = ChatPromptTemplate.from_template(
"""
你需要把用户当前问题改写成一个独立、完整、适合检索资料的问题。
要求:
1. 只输出改写后的问题,不要解释。
2. 不要回答问题。
3. 不要编造历史对话和用户问题里没有的信息。
4. 如果问题有歧义,不要强行缩窄范围,要保留可能含义。
5. 如果当前问题本身已经完整,就原样输出。
例子:
用户问:发育缓慢怎么办
较好改写:儿童或青少年发育缓慢怎么办,包括身高体重增长慢、生长迟缓等情况的饮食、运动、睡眠和就医建议
历史对话:
{chat_history}
用户当前问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
rewritten_question = chain.invoke({
"chat_history": chat_history,
"user_question": user_question,
}).strip()
print(f"原问题:{user_question}")
print(f"改写后问题:{rewritten_question}")
return rewritten_question
def generate_search_queries(
user_question: str,
rewritten_question: str,
llm,
) -> List[str]:
"""
✅ 优化点:生成多个检索词,提升召回率。
召回率说人话:别漏掉可能有用的资料。
"""
prompt = ChatPromptTemplate.from_template(
"""
你是 RAG 检索词生成器。
请根据用户问题,生成 2 到 4 个适合检索资料的中文查询语句。
要求:
1. 每行一个查询语句。
2. 不要解释。
3. 不要回答问题。
4. 不要编造具体年龄、疾病诊断、检查结果。
5. 如果问题有歧义,要保留多个可能方向。
6. 查询语句要短一点,适合搜索资料。
用户原问题:
{user_question}
改写后问题:
{rewritten_question}
"""
)
chain = prompt | llm | StrOutputParser()
raw_text = chain.invoke({
"user_question": user_question,
"rewritten_question": rewritten_question,
}).strip()
queries = [user_question, rewritten_question]
for line in raw_text.splitlines():
line = line.strip()
line = re.sub(r"^[\-\*\d\.\、\)\s]+", "", line).strip()
if line:
queries.append(line)
# ✅ 优化点:去重,保留顺序
final_queries = []
seen = set()
for q in queries:
if q not in seen:
final_queries.append(q)
seen.add(q)
print("检索查询:")
for q in final_queries:
print(f"- {q}")
return final_queries[:5]
# =========================
# 7. BM25 相关
# =========================
def chinese_tokenizer(text: str) -> List[str]:
"""
给 BM25 用的中文分词函数。
分词说人话:把一句中文切成一个个词。
"""
return [word.strip() for word in jieba.cut(text) if word.strip()]
def save_bm25_docs(chunks: List[Document]):
"""
保存 chunks,给 BM25 使用。
"""
with open(BM25_DOCS_PATH, "wb") as f:
pickle.dump(chunks, f)
print(f"BM25 文档已保存到 {BM25_DOCS_PATH}")
def load_bm25_retriever(k: int = 5):
"""
加载 BM25 检索器。
BM25 说人话:按关键词找资料。
"""
if not os.path.exists(BM25_DOCS_PATH):
raise FileNotFoundError(
f"找不到 {BM25_DOCS_PATH},请先选择入库模式。"
)
with open(BM25_DOCS_PATH, "rb") as f:
docs = pickle.load(f)
retriever = BM25Retriever.from_documents(
docs,
preprocess_func=chinese_tokenizer,
)
retriever.k = k
return retriever
# =========================
# 8. 入库、删除
# =========================
def ingest_documents(vector_store, chunks: List[Document]):
"""
入库:
1. 存进 Chroma,给向量检索用。
2. 存成 pkl,给 BM25 用。
"""
# ✅ 优化点:使用稳定 ID,避免 text_new_chunk_0 这种每次都重复、含义不清的 ID
ids = [stable_doc_id(doc, i) for i, doc in enumerate(chunks)]
vector_store.add_documents(
documents=chunks,
ids=ids,
)
save_bm25_docs(chunks)
print(f"成功入库 {len(chunks)} 个 chunk")
def delete_collection(vector_store):
"""
删除 Chroma collection 和 BM25 文件。
collection 说人话:知识库里的一张表。
"""
try:
vector_store.delete_collection()
print(f"{COLLECTION_NAME} 已删除")
except Exception as e:
print(f"删除 Chroma collection 时出现问题:{e}")
if os.path.exists(BM25_DOCS_PATH):
os.remove(BM25_DOCS_PATH)
print(f"{BM25_DOCS_PATH} 已删除")
# =========================
# 9. 混合检索
# =========================
def dedupe_documents(docs: List[Document]) -> List[Document]:
"""
文档去重。
"""
result = []
seen_keys = set()
for doc in docs:
key = (
doc.metadata.get("source", ""),
doc.metadata.get("page", ""),
doc.page_content[:100],
)
if key in seen_keys:
continue
if is_low_quality_text(doc.page_content):
continue
seen_keys.add(key)
result.append(doc)
return result
def hybrid_search(
search_queries: List[str],
rerank_query: str,
bm25_retriever,
vector_store,
reranker,
bm25_k: int = 10,
vector_k: int = 10,
final_k: int = 5,
) -> List[Document]:
"""
混合检索 + rerank:
1. BM25 按关键词找
2. Chroma 按意思找
3. 多查询合并
4. 去重
5. 过滤低质量片段
6. rerank 重排
"""
all_docs = []
for query in search_queries:
# 1. 关键词检索
bm25_retriever.k = bm25_k
keyword_docs = bm25_retriever.invoke(query)
# 2. 语义检索
semantic_docs = vector_store.similarity_search(
query,
k=vector_k,
)
all_docs.extend(keyword_docs)
all_docs.extend(semantic_docs)
# ✅ 优化点:多路检索后统一去重和过滤
candidate_docs = dedupe_documents(all_docs)
if not candidate_docs:
return []
# 3. rerank
reranked_docs = reranker.compress_documents(
documents=candidate_docs,
query=rerank_query,
)
# ✅ 优化点:rerank 后再过滤一次,防止目录类片段混进来
reranked_docs = [
doc for doc in reranked_docs
if not is_low_quality_text(doc.page_content)
]
return list(reranked_docs)[:final_k]
# =========================
# 10. 构造上下文
# =========================
def build_context(docs: List[Document]) -> str:
"""
把检索到的资料整理成 prompt 里的参考文档。
"""
context_list = []
for i, doc in enumerate(docs, start=1):
source = doc.metadata.get("source", "未知来源")
page = page_show(doc)
context_list.append(
f"""
[资料{i}]
来源:{source}
页码:第 {page} 页
内容:
{doc.page_content}
"""
)
return "\n\n".join(context_list)
def build_source_text(docs: List[Document]) -> str:
"""
构造展示给用户看的参考资料。
"""
return "\n\n".join([
f"-------------------- 资料 {i + 1} --------------------\n"
f"页码:第 {page_show(doc)} 页\n"
f"来源:{doc.metadata.get('source', '未知来源')}\n\n"
f"{doc.page_content[:1000]}"
for i, doc in enumerate(docs)
])
# =========================
# 11. 问答
# =========================
def answer_question(
user_question: str,
llm,
bm25_retriever,
vector_store,
reranker,
memory: SimpleWindowMemory,
) -> str:
"""
完整问答流程:
1. 查询重写
2. 多检索词生成
3. 混合检索
4. rerank
5. 拼接上下文
6. 调用大模型
7. 保存本轮问答到记忆
8. 返回答案和来源
"""
rewritten_question = rewrite_query(
user_question=user_question,
llm=llm,
memory=memory,
)
search_queries = generate_search_queries(
user_question=user_question,
rewritten_question=rewritten_question,
llm=llm,
)
docs = hybrid_search(
search_queries=search_queries,
rerank_query=rewritten_question,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
bm25_k=BM25_K,
vector_k=VECTOR_K,
final_k=FINAL_K,
reranker=reranker,
)
# ✅ 优化点:没搜到资料时,别让模型硬答
if not docs:
answer = "文档资料没有提到,我不知道。建议补充更多相关权威资料后再查询。"
memory.save(user_question, answer)
return (
"\n==================== 回答 ====================\n\n"
f"{answer}"
"\n\n==================== 参考资料 ====================\n\n"
"无"
)
context = build_context(docs)
chat_history = memory.load()
# ✅ 优化点:prompt 更严格,尤其是健康类问题不要乱补
prompt = ChatPromptTemplate.from_template(
"""
你是一个健康资料问答助手,但不是医生。
你只能根据[参考文档]回答用户问题。
历史对话只用于理解用户当前问题,不能当成资料来源。
重要要求:
1. 每一条关键建议都必须来自参考文档。
2. 每一条关键建议后面都要标注来源编号,例如:[资料1]、[资料2]。
3. 如果参考文档里没有依据,就不要说。
4. 如果参考文档里没有答案,直接回答:文档资料没有提到,我不知道。
5. 不要使用你自己的医学常识补充答案。
6. 不要编造参考文档里没有的内容。
7. 不要把“生长迟缓”和所有“发育迟缓”混为一谈。
8. 如果资料主要讲的是身高、体重、生长迟缓,就要说明回答范围主要是这一类。
9. 如果参考文档提到需要就医,要把就医提醒放在前面。
10. 回答要通俗易懂,说人话。
11. 你只是提供健康建议,不能代替医生诊断。
回答格式:
先用 1 到 2 句话说结论。
然后分条回答。
最后提醒不能代替医生诊断。
历史对话:
{chat_history}
参考文档:
{context}
用户问题:
{user_question}
"""
)
chain = prompt | llm | StrOutputParser()
answer = chain.invoke({
"chat_history": chat_history,
"context": context,
"user_question": user_question,
}).strip()
memory.save(user_question, answer)
source_text = build_source_text(docs)
return (
"\n==================== 回答 ====================\n\n"
f"{answer}"
"\n\n==================== 参考资料 ====================\n\n"
f"{source_text}"
)
# =========================
# 12. 菜单功能
# =========================
def run_ingest(vector_store):
"""
入库模式。
"""
raw_docs = load_pdf(PDF_PATH)
chunks = split_documents(
raw_docs,
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
)
if not chunks:
print("没有有效 chunk,入库失败。")
return
ingest_documents(vector_store, chunks)
def run_qa(llm, vector_store, reranker, memory):
"""
提问模式。
"""
bm25_retriever = load_bm25_retriever(k=BM25_K)
while True:
user_question = input("你想了解什么(输入 q 退出,输入 clear 清空记忆):").strip()
if user_question.lower() == "q":
print("已退出提问模式")
break
if user_question.lower() == "clear":
memory.clear()
print("已清空上下文记忆")
continue
if not user_question:
print("问题不能为空")
continue
try:
answer = answer_question(
user_question=user_question,
llm=llm,
bm25_retriever=bm25_retriever,
vector_store=vector_store,
reranker=reranker,
memory=memory,
)
print("\n" + answer + "\n")
# ✅ 优化点:问答出错时不让整个程序直接崩
except Exception as e:
print(f"问答过程出错:{e}")
def run_delete(vector_store):
"""
删除模式。
"""
confirm = input(f"确认删除 {COLLECTION_NAME} 吗?输入 yes:").strip()
if confirm == "yes":
delete_collection(vector_store)
else:
print("已取消删除")
# =========================
# 13. 主函数
# =========================
def main():
"""
主函数只做三件事:
1. 初始化模型和数据库
2. 初始化记忆
3. 根据用户选择调用不同功能
"""
embedding = create_embedding_model()
vector_store = create_vector_store(embedding)
llm = create_llm()
reranker = create_reranker()
memory = create_memory()
while True:
print("\n请选择模式:")
print("1 = 入库")
print("2 = 提问")
print("3 = 删除 collection")
print("q = 退出")
mode = input("请输入:").strip()
if mode == "1":
try:
run_ingest(vector_store)
except Exception as e:
print(f"入库失败:{e}")
elif mode == "2":
try:
run_qa(llm, vector_store, reranker, memory)
except Exception as e:
print(f"提问模式启动失败:{e}")
elif mode == "3":
run_delete(vector_store)
elif mode.lower() == "q":
print("程序已退出")
break
else:
print("请输入 1、2、3 或 q")
if __name__ == "__main__":
main()
Tips:同样本篇也为YellowDuck小白制作,如有问题可以指出,YellowDuck尽快修改!!!
本文来自博客园,作者:LiYellowDuck,转载请注明原文链接:https://www.cnblogs.com/LiYellowDuck/p/20008044,YellowDuck热爱手搓!!!

浙公网安备 33010602011771号