RAG优化(续一)

这里是为了续上上一章节的从RAG入门到RAG优化,所以不多废话,直接上手

RAG优化

(1)chunk优化- [YES]

(2)BM25关键检索- [YES]

(3)rerank 重排

其实很好理解,rerank的实现过程就是首先粗略的找出一些相关的,然后再去使用更聪明的模型去进行排名,然后选出最有用的放在前面,为什么需要重排序呢,其实也很好理解,第一次筛是大范围粗略的筛,所以重排就需要更加的精细,使得检索更加准确

这里是AI生成的流程图,以便于理解整个流程

ChatGPT Image 2026年5月10日 16_58_33

之所以放上流程图是因为加上许多优化,我们的代码已经比较复杂来到了大概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尽快修改!!!

posted @ 2026-05-12 19:03  LiYellowDuck  阅读(23)  评论(0)    收藏  举报