import os import uuid from typing import List from pydantic import Field from sentence_transformers import CrossEncoder from langchain_openai import OpenAIEmbeddings, ChatOpenAI from langchain_chroma import Chroma from langchain_text_splitters import RecursiveCharacterTextSplitter from langchain_core.prompts import ChatPromptTemplate from langchain_core.runnables import RunnablePassthrough from langchain_core.output_parsers import StrOutputParser from langchain_core.documents import Document from langchain_classic.retrievers.multi_query import MultiQueryRetriever from langchain_classic.retrievers import ContextualCompressionRetriever from langchain_classic.retrievers.document_compressors import LLMChainFilter from langchain_core.retrievers import BaseRetriever from langchain_core.documents import Document
CHUNK_SIZE = 300 CHUNK_OVERLAP = 50 CHROMA_PERSIST_DIR = "./advanced_rag_chroma" RETRIEVE_TOP_K = 4 FETCH_K = 10 LAMBDA_MULT = 0.3
llm = ChatOpenAI( model="qwen2.5-1.5b-instruct-q4_k_m.gguf", base_url="http://127.0.0.1:11433/v1", api_key="dummy", temperature=0.3, max_tokens=800, )
embeddings = OpenAIEmbeddings( model="qwen3-embedding-local.gguf", base_url="http://127.0.0.1:11434/v1", api_key="dummy" )
text_splitter = RecursiveCharacterTextSplitter( chunk_size=CHUNK_SIZE, chunk_overlap=CHUNK_OVERLAP, separators=["\n\n", "\n", "。", ",", " "] )
def add_documents_safe(db, docs): ids = [str(uuid.uuid4()) for _ in docs] db.add_documents(docs, ids=ids) print(f"[+] 本次追加 {len(docs)} 个文本块,当前总数量:{db._collection.count()}")
def get_vector_store(documents: List[Document], incremental: bool = True) -> Chroma: if not incremental: if os.path.exists(CHROMA_PERSIST_DIR): import shutil shutil.rmtree(CHROMA_PERSIST_DIR) print("[*] 删除旧向量库,覆盖重建模式")
split_docs = text_splitter.split_documents(documents) if os.path.exists(CHROMA_PERSIST_DIR): print("[+] 向量库已存在,增量追加") db = Chroma(persist_directory=CHROMA_PERSIST_DIR, embedding_function=embeddings) add_documents_safe(db, split_docs) else: print("[*] 新建向量库") db = Chroma.from_documents( documents=split_docs, embedding=embeddings, persist_directory=CHROMA_PERSIST_DIR ) print(f"[+] 存入 {len(split_docs)} 个文本块") return db
def format_docs(docs: List[Document]) -> str: return "\n---\n".join( f"[来源:{doc.metadata.get('source','未知')}|页码:{doc.metadata.get('page','-')}]\n{doc.page_content}" for doc in docs )
def build_advanced_retriever(vector_db: Chroma): base_retriever = vector_db.as_retriever( search_type="mmr", search_kwargs={ "k": RETRIEVE_TOP_K, "fetch_k": FETCH_K, "lambda_mult": LAMBDA_MULT } )
multi_query_prompt = ChatPromptTemplate.from_messages([ ("system", """你是查询生成助手。针对用户问题,生成3个不同角度、不同措辞的检索查询,用于知识库向量检索。只输出查询,每行一条,不要多余解释。"""), ("human", "原始问题:{question}") ])
multi_query_retriever = MultiQueryRetriever.from_llm( retriever=base_retriever, llm=llm, prompt=multi_query_prompt )
compressor = LLMChainFilter.from_llm(llm) compression_retriever = ContextualCompressionRetriever( base_retriever=multi_query_retriever, base_compressor=compressor ) return compression_retriever
def build_advanced_rag_chain(vector_db: Chroma): retriever = build_advanced_retriever(vector_db)
rag_prompt = ChatPromptTemplate.from_messages([ ("system", """你是企业知识库问答助手,严格依据提供的上下文回答。 1. 只使用上下文给出的信息,不要编造;知识库没有则输出“知识库中未找到相关内容”。 2. 回答尽量简洁准确,可以引用来源信息。 上下文参考: {context}"""), ("human", "{question}") ])
advanced_rag_chain = ( {"context": retriever | format_docs, "question": RunnablePassthrough()} | rag_prompt | llm | StrOutputParser() ) return advanced_rag_chain, retriever
class RerankerRetriever(BaseRetriever): base_retriever: BaseRetriever = Field(description="底层召回检索器") reranker_model: CrossEncoder = Field(description="交叉编码器重排模型") top_n: int = Field(default=3, description="重排之后保留多少条")
def _get_relevant_documents(self, query: str) -> List[Document]: candidates = self.base_retriever.invoke(query) if not candidates: return []
pairs = [[query, doc.page_content] for doc in candidates] scores = self.reranker_model.predict(pairs) scored_docs = sorted(zip(candidates, scores), key=lambda x: x[1], reverse=True) keep_docs = [doc for doc, score in scored_docs[:self.top_n]]
print(f"\n[Reranker重排后保留 {len(keep_docs)} 条文档]") for d, s in scored_docs[:self.top_n]: print(f"rerank_score={s:.4f} | source={d.metadata.get('source')}") return keep_docs
def build_reranker_advanced_rag(vector_db: Chroma): base_compress_retriever = build_advanced_retriever(vector_db) reranker = CrossEncoder("./models/models/BAAI--bge-reranker-v2-m3/snapshots/master") rerank_retriever = RerankerRetriever( base_retriever=base_compress_retriever, reranker_model=reranker, top_n=3 )
rag_prompt = ChatPromptTemplate.from_messages([ ("system", """你是知识库问答助手,严格依据提供的上下文回答。无相关信息直接输出“知识库中未找到相关内容”,禁止幻觉编造。上下文:{context}"""), ("human", "{question}") ]) chain = ( {"context": rerank_retriever | format_docs, "question": RunnablePassthrough()} | rag_prompt | llm | StrOutputParser() ) return chain, rerank_retriever
if __name__ == "__main__": test_docs = [ Document( page_content="llama.cpp是高性能GGUF格式本地大模型推理框架,支持CPU/GPU混合加速,能够对外提供兼容OpenAI接口的本地API服务,本程序使用该服务提供LLM能力。", metadata={"source":"local_env.md"} ), Document( page_content="Embedding向量不能直接使用对话型大模型,必须使用专门的嵌入模型;本示例使用Qwen3‑Embedding作为向量模型,用于Chroma向量库的文档向量化检索。", metadata={"source":"embedding_note.md"} ), Document( page_content="Naive‑RAG即朴素RAG,仅做基础向量相似度检索,没有多查询扩展、没有上下文过滤压缩、没有重排序模块,检索角度单一,复杂问题召回效果有限。", metadata={"source":"rag_compare.md"} ), Document( page_content="Advanced‑RAG在Naive‑RAG朴素RAG基础上做能力增强,典型优化手段包含MultiQuery多查询生成、MMR多样性检索、LLM上下文压缩过滤、Cross‑Encoder重排序。本程序完整实现以上增强链路。", metadata={"source":"rag_compare.md"} ), Document( page_content="Modular‑RAG将RAG流程拆成可插拔组件,检索器、压缩器、重排器都可以自由替换,但整体执行流程是固定的,本示例的链路就是模块化RAG的实践。", metadata={"source":"rag_intro.md"} ), Document( page_content="Agentic‑RAG依靠大模型自主规划决策,动态判断是否需要多次检索,适合多跳、复杂推理类问题,和本示例Advanced‑RAG固定检索链路有明显区别。", metadata={"source":"rag_intro.md"} ), ]
db = get_vector_store(test_docs, incremental=False) print(f"\n向量库总块数:{db._collection.count()}")
rag_chain, ret = build_reranker_advanced_rag(db)
user_query = "Advanced‑RAG相比Naive‑RAG做了哪些增强手段?" print(f"\n[用户问题:{user_query}]")
retrieved_docs = ret.invoke(user_query) print("\n[经过多查询+压缩+重排之后的上下文]") print(format_docs(retrieved_docs))
answer = rag_chain.invoke(user_query) print("\n[Advanced‑RAG最终回答]") print(answer)
|