""" RAG 全链路优化系统 优化点: 1. 语义切块(SemanticChunker)替代固定切片 2. 多路召回(BM25 + 向量检索 + 混合检索) 3. 查询侧优化(查询重写、查询分解、HyDE) 4. Reranker 重排序精排 5. Milvus 向量数据库替代 Chroma """ import os import re import json from typing import List, Optional from concurrent.futures import ThreadPoolExecutor from dotenv import load_dotenv # LangChain 核心组件 from langchain_community.document_loaders import PyMuPDFLoader from langchain_text_splitters import RecursiveCharacterTextSplitter from langchain_community.embeddings import DashScopeEmbeddings from langchain_community.vectorstores import Milvus from langchain_community.retrievers import BM25Retriever from langchain_classic.retrievers import EnsembleRetriever from langchain_classic.retrievers.multi_query import MultiQueryRetriever from langchain_core.prompts import ChatPromptTemplate, PromptTemplate from langchain_openai import ChatOpenAI from langchain_core.output_parsers import StrOutputParser from langchain_core.documents import Document from langchain_experimental.text_splitter import SemanticChunker # 加载环境变量 load_dotenv() # ==================================================================== # 日志工具 # ==================================================================== def log_info(message: str): """打印日志信息""" print(f"[INFO] {message}") def log_error(message: str): """打印错误日志""" print(f"[ERROR] {message}") def log_warn(message: str): """打印警告日志""" print(f"[WARN] {message}") # ==================================================================== # 查询侧优化模块 # ==================================================================== class QueryOptimizer: """查询侧优化器:查询重写、查询分解、HyDE""" def __init__(self, llm: ChatOpenAI): self.llm = llm self._init_chains() def _init_chains(self): """初始化各优化链""" # 查询重写 Prompt self.rewrite_prompt = PromptTemplate( input_variables=["query"], template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。 改写要求: 1. 补充隐含的上下文信息 2. 将口语化表达转为专业表述 3. 消除歧义,明确查询意图 4. 保持原意不变,不要添加原问题未提及的内容 5. 直接输出改写后的查询,不要解释 用户问题: {query} 改写后的查询:""" ) self.rewrite_chain = self.rewrite_prompt | self.llm | StrOutputParser() # 查询分解 Prompt self.decompose_prompt = PromptTemplate( input_variables=["question"], template="""你是一个问题分解助手。请将用户的复杂问题分解为 2-4 个独立的子问题, 每个子问题应该能独立检索和回答。 要求: 1. 子问题之间互不依赖,可以并行检索 2. 子问题覆盖原始问题的所有方面 3. 每个子问题简洁明确 4. 以 JSON 数组格式输出 用户问题: {question} 输出格式: ["子问题1", "子问题2", "子问题3"] 子问题列表:""" ) self.decompose_chain = self.decompose_prompt | self.llm | StrOutputParser() # HyDE Prompt(假设性文档嵌入) self.hyde_prompt = PromptTemplate( input_variables=["question"], template="""请根据以下问题,写一段可能出现在相关文档中的内容(约100-200字)。 不需要完全准确,只需要包含相关的关键词和表述方式。 问题: {question} 假设性文档内容:""" ) self.hyde_chain = self.hyde_prompt | self.llm | StrOutputParser() def rewrite(self, query: str) -> str: """查询重写:将口语化问题转为精确检索查询""" try: rewritten = self.rewrite_chain.invoke({"query": query}) log_info(f"查询重写: '{query}' → '{rewritten}'") return rewritten except Exception as e: log_warn(f"查询重写失败: {e},使用原始查询") return query def decompose(self, question: str) -> List[str]: """查询分解:将复杂问题拆分为子问题""" try: result = self.decompose_chain.invoke({"question": question}) sub_queries = json.loads(result.strip()) log_info(f"查询分解: '{question}' → {len(sub_queries)} 个子问题") return sub_queries except (json.JSONDecodeError, Exception) as e: log_warn(f"查询分解失败: {e},使用原始问题") return [question] def hyde(self, question: str) -> str: """HyDE:生成假设性文档用于检索""" try: hypothetical_doc = self.hyde_chain.invoke({"question": question}) log_info(f"HyDE 生成假设文档: {hypothetical_doc[:80]}...") return hypothetical_doc except Exception as e: log_warn(f"HyDE 生成失败: {e},使用原始问题") return question # ==================================================================== # Reranker 重排序模块 # ==================================================================== class LLMReranker: """基于 LLM 的重排序器(无需额外模型,用大模型做相关性打分)""" def __init__(self, llm: ChatOpenAI): self.llm = llm self.score_prompt = PromptTemplate( input_variables=["query", "document"], template="""请评估以下文档与用户问题的相关性,给出 0-10 的分数。 只输出数字分数,不要任何解释。 用户问题: {query} 文档内容: {document} 相关性分数(0-10):""" ) self.score_chain = self.score_prompt | self.llm | StrOutputParser() def rerank(self, query: str, documents: List[Document], top_k: int = 3) -> List[Document]: """对检索结果重排序,返回最相关的 top_k 个文档""" if not documents: return [] if len(documents) <= top_k: return documents log_info(f"Reranker: 对 {len(documents)} 个文档重排序,保留 Top-{top_k}") scored_docs = [] for doc in documents: try: score_str = self.score_chain.invoke({ "query": query, "document": doc.page_content[:500] }) # 提取数字分数 score = float(re.search(r'[\d.]+', score_str).group()) scored_docs.append((doc, score)) except Exception: # 打分失败时给默认分数 scored_docs.append((doc, 5.0)) # 按分数降序排序 scored_docs.sort(key=lambda x: x[1], reverse=True) # 返回 top_k 个 result = [doc for doc, score in scored_docs[:top_k]] log_info(f"Reranker 完成,分数范围: {scored_docs[0][1]:.1f} - {scored_docs[-1][1]:.1f}") return result # ==================================================================== # 核心 RAG 系统(全链路优化版) # ==================================================================== class OptimizedRAGSystem: """ 全链路优化 RAG 系统 优化链路: 文档加载 → 数据清洗 → 语义切块 → Milvus存储 → 多路召回 → Reranker精排 → 生成回答 查询优化: 用户问题 → 查询重写/分解/HyDE → 多路召回 → 重排序 → 生成 """ def __init__( self, pdf_path: str, collection_name: str = "rag_optimized", milvus_uri: str = "http://localhost:19530", use_semantic_chunking: bool = True, ): """ 初始化优化版 RAG 系统 Args: pdf_path: PDF 文件路径 collection_name: Milvus 集合名称 milvus_uri: Milvus 连接地址 use_semantic_chunking: 是否使用语义切块(需要较多 API 调用) """ self.pdf_path = pdf_path self.collection_name = collection_name self.milvus_uri = milvus_uri self.use_semantic_chunking = use_semantic_chunking self.pages = None self.split_docs = None self.vectorstore = None # 初始化 Embedding 模型 self.embedding_model = DashScopeEmbeddings( model=os.getenv("EMBEDDING_MODEL", "text-embedding-v3"), dashscope_api_key=os.getenv("DASHSCOPE_API_KEY", "") ) # 初始化大模型 self.llm = ChatOpenAI( model=os.getenv("MODEL_NAME", "deepseek-chat"), openai_api_key=os.getenv("OPENAI_API_KEY", ""), openai_api_base=os.getenv("OPENAI_API_BASE", "https://api.deepseek.com/v1"), temperature=0.7 ) # 初始化查询优化器和重排序器 self.query_optimizer = QueryOptimizer(self.llm) self.reranker = LLMReranker(self.llm) log_info("全链路优化 RAG 系统初始化完成") log_info(f"模型: {os.getenv('MODEL_NAME', 'deepseek-chat')}") log_info(f"Milvus: {milvus_uri}") log_info(f"切块策略: {'语义切块' if use_semantic_chunking else '递归字符切块'}") # ================================================================ # 文档处理阶段 # ================================================================ def load_pdf(self) -> List[Document]: """加载 PDF 文档""" log_info(f"加载 PDF: {self.pdf_path}") if not os.path.exists(self.pdf_path): raise FileNotFoundError(f"PDF 文件不存在: {self.pdf_path}") loader = PyMuPDFLoader(self.pdf_path) self.pages = loader.load() total_content = sum(len(p.page_content) for p in self.pages) if total_content == 0: raise ValueError("PDF 文件无可提取文本内容,请使用文本型 PDF 或进行 OCR 处理") log_info(f"加载完成: {len(self.pages)} 页, {total_content} 字符") return self.pages def clean_text(self, text: str) -> str: """清洗文本:去除噪声、规范化""" if not text or not text.strip(): return "" # 删除多余换行(保留最多2个) text = re.sub(r'\n{3,}', '\n\n', text) # 删除项目符号 text = text.replace('•', '').replace('·', '') # 合并多余空格 text = re.sub(r'[^\S\n]+', ' ', text) # 删除控制字符 text = re.sub(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f-\x9f]', '', text) return text.strip() def clean_all_pages(self) -> List[Document]: """清洗所有页面""" if not self.pages: raise ValueError("请先调用 load_pdf()") log_info("清洗文档...") for page in self.pages: page.page_content = self.clean_text(page.page_content) # 过滤空页 self.pages = [p for p in self.pages if p.page_content.strip()] log_info(f"清洗完成: {len(self.pages)} 页有效内容") return self.pages def split_documents(self) -> List[Document]: """ 文档切块(支持语义切块和递归字符切块两种策略) 语义切块:基于相邻句子的语义相似度,在话题转变处切分 递归字符切块:按字符数切分,带重叠防止上下文断裂 """ if not self.pages: raise ValueError("请先调用 load_pdf()") if self.use_semantic_chunking: log_info("使用语义切块策略(SemanticChunker)...") try: text_splitter = SemanticChunker( embeddings=self.embedding_model, breakpoint_threshold_type="percentile", # 90 表示只有相似度排名后 10% 的位置才切分 # 值越大切得越少,块越大 breakpoint_threshold_amount=90, ) self.split_docs = text_splitter.split_documents(self.pages) except Exception as e: log_warn(f"语义切块失败: {e},回退到递归字符切块") self.split_docs = self._fallback_split() else: self.split_docs = self._fallback_split() # 过滤空块 self.split_docs = [ doc for doc in self.split_docs if doc.page_content and doc.page_content.strip() ] log_info(f"切块完成: {len(self.split_docs)} 个文档块") if self.split_docs: avg_len = sum(len(d.page_content) for d in self.split_docs) / len(self.split_docs) log_info(f"平均块长度: {avg_len:.0f} 字符") return self.split_docs def _fallback_split(self) -> List[Document]: """回退策略:递归字符切块""" log_info("使用递归字符切块策略...") text_splitter = RecursiveCharacterTextSplitter( separators=["\n\n", "\n", "。", "!", "?", ";", " ", ""], chunk_size=300, chunk_overlap=50, length_function=len, ) return text_splitter.split_documents(self.pages) # ================================================================ # 向量存储与多路召回 # ================================================================ def create_vectorstore(self) -> Milvus: """将文档块存入 Milvus 向量数据库""" if not self.split_docs: raise ValueError("请先调用 split_documents()") log_info(f"创建 Milvus 向量库: {self.collection_name}") self.vectorstore = Milvus.from_documents( documents=self.split_docs, embedding=self.embedding_model, connection_args={"uri": self.milvus_uri}, collection_name=self.collection_name, drop_old=True, # 开发阶段覆盖旧数据 ) log_info(f"Milvus 集合创建完成: {len(self.split_docs)} 个文档块已入库") return self.vectorstore def load_vectorstore(self) -> Milvus: """加载已有的 Milvus 向量数据库""" log_info(f"加载 Milvus 集合: {self.collection_name}") self.vectorstore = Milvus( embedding_function=self.embedding_model, connection_args={"uri": self.milvus_uri}, collection_name=self.collection_name, ) log_info("Milvus 向量库加载完成") return self.vectorstore def build_ensemble_retriever(self, k: int = 10) -> EnsembleRetriever: """ 构建混合检索器(BM25 + 向量检索) BM25: 关键词精确匹配,擅长处理精确术语 向量检索: 语义相似度匹配,擅长处理同义词和语义理解 混合: 两者加权融合,取长补短 """ if not self.vectorstore or not self.split_docs: raise ValueError("请先创建向量库") # BM25 检索器(关键词匹配) bm25_retriever = BM25Retriever.from_documents(self.split_docs) bm25_retriever.k = k # 向量检索器(语义匹配) vector_retriever = self.vectorstore.as_retriever( search_kwargs={"k": k} ) # 混合检索器:BM25 权重 0.4,向量检索权重 0.6 ensemble_retriever = EnsembleRetriever( retrievers=[bm25_retriever, vector_retriever], weights=[0.4, 0.6], ) log_info(f"混合检索器构建完成 (BM25:0.4 + Vector:0.6, Top-{k})") return ensemble_retriever def build_multi_query_retriever( self, base_retriever: EnsembleRetriever ) -> MultiQueryRetriever: """ 构建多查询检索器 原理:用 LLM 将一个问题改写为多个不同角度的问题, 分别检索后合并去重,增加召回率 """ multi_query_prompt = PromptTemplate( input_variables=["question"], template="""你是一个问题改写助手。请为下面的问题生成 3 个不同的改写版本。 每个版本应从不同角度表达相同意图,使用不同的关键词。 每个问题单独一行,不要编号。 原始问题: {question} 改写后的问题:""" ) retriever = MultiQueryRetriever.from_llm( retriever=base_retriever, llm=self.llm, prompt=multi_query_prompt, ) log_info("多查询检索器构建完成") return retriever # ================================================================ # 检索与生成 # ================================================================ def retrieve_with_optimization( self, query: str, retriever, mode: str = "rewrite", rerank_top_k: int = 3, ) -> List[Document]: """ 带查询优化 + Reranker 的检索流程 Args: query: 用户原始问题 retriever: 检索器实例 mode: 查询优化模式 - "none": 不做优化,直接检索 - "rewrite": 查询重写后检索 - "decompose": 查询分解后并行检索 - "hyde": HyDE 假设文档检索 - "multi_query": 多查询检索(需要传入 MultiQueryRetriever) rerank_top_k: Reranker 精排后保留的文档数 Returns: 精排后的文档列表 """ log_info(f"检索模式: {mode} | 问题: {query}") # 第一步:查询优化 if mode == "rewrite": optimized_query = self.query_optimizer.rewrite(query) raw_docs = retriever.invoke(optimized_query) elif mode == "decompose": sub_queries = self.query_optimizer.decompose(query) raw_docs = self._parallel_retrieve(sub_queries, retriever) elif mode == "hyde": hypothetical_doc = self.query_optimizer.hyde(query) raw_docs = retriever.invoke(hypothetical_doc) elif mode == "multi_query": # MultiQueryRetriever 内部自带多查询逻辑 raw_docs = retriever.invoke(query) else: # mode == "none" raw_docs = retriever.invoke(query) log_info(f"初步召回: {len(raw_docs)} 个文档") # 第二步:Reranker 精排 if len(raw_docs) > rerank_top_k: reranked_docs = self.reranker.rerank(query, raw_docs, top_k=rerank_top_k) else: reranked_docs = raw_docs log_info(f"精排后: {len(reranked_docs)} 个文档") return reranked_docs def _parallel_retrieve( self, queries: List[str], retriever ) -> List[Document]: """并行检索多个查询并合并去重""" all_docs = [] seen_contents = set() def retrieve_one(q): return retriever.invoke(q) with ThreadPoolExecutor(max_workers=min(len(queries), 4)) as executor: futures = [executor.submit(retrieve_one, q) for q in queries] for future in futures: try: docs = future.result() for doc in docs: content_hash = hash(doc.page_content) if content_hash not in seen_contents: seen_contents.add(content_hash) all_docs.append(doc) except Exception as e: log_warn(f"并行检索子任务失败: {e}") return all_docs def generate_answer(self, query: str, relevant_docs: List[Document]) -> str: """基于检索结果生成回答""" if not relevant_docs: return "抱歉,根据现有资料,我找不到与您问题相关的信息。" log_info("生成回答...") context = "\n\n---\n\n".join([doc.page_content for doc in relevant_docs]) prompt = ChatPromptTemplate.from_template("""你是一个专业的知识库助手。请根据以下检索到的上下文回答用户问题。 **规则:** - 只基于提供的上下文回答,不要编造信息 - 如果上下文中没有相关信息,请明确说明 - 回答要简洁直接,条理清晰 - 引用原文时用引号标注 **检索到的上下文:** {context} **用户问题:** {question} **回答:**""") chain = prompt | self.llm | StrOutputParser() answer = chain.invoke({"context": context, "question": query}) log_info("回答生成完成") return answer def query( self, question: str, mode: str = "rewrite", rerank_top_k: int = 3, ) -> dict: """ 完整查询流程:查询优化 → 多路召回 → 重排序 → 生成回答 Args: question: 用户问题 mode: 查询优化模式 ("none"/"rewrite"/"decompose"/"hyde"/"multi_query") rerank_top_k: 精排后保留文档数 Returns: 包含回答和中间结果的字典 """ # 构建检索器 ensemble_retriever = self.build_ensemble_retriever(k=10) # 如果是 multi_query 模式,额外包装一层 if mode == "multi_query": retriever = self.build_multi_query_retriever(ensemble_retriever) else: retriever = ensemble_retriever # 检索 + 重排序 relevant_docs = self.retrieve_with_optimization( query=question, retriever=retriever, mode=mode, rerank_top_k=rerank_top_k, ) # 生成回答 answer = self.generate_answer(question, relevant_docs) return { "question": question, "mode": mode, "doc_count": len(relevant_docs), "answer": answer, "sources": [doc.page_content[:100] for doc in relevant_docs], } # ================================================================ # 构建与对比 # ================================================================ def build_knowledge_base(self): """构建完整知识库:加载 → 清洗 → 切块 → 向量化""" log_info("=" * 60) log_info("开始构建优化版知识库...") log_info("=" * 60) self.load_pdf() self.clean_all_pages() self.split_documents() self.create_vectorstore() log_info("=" * 60) log_info("知识库构建完成!") log_info("=" * 60) def compare_retrieval_modes(self, question: str): """对比不同检索模式的效果""" log_info("=" * 60) log_info(f"对比测试: {question}") log_info("=" * 60) modes = ["none", "rewrite", "decompose", "hyde"] results = {} for mode in modes: log_info(f"\n--- 模式: {mode} ---") try: result = self.query(question, mode=mode) results[mode] = result print(f"\n[{mode}] 召回 {result['doc_count']} 个文档") print(f"回答: {result['answer'][:200]}...") except Exception as e: log_error(f"模式 {mode} 执行失败: {e}") results[mode] = {"error": str(e)} return results # ==================================================================== # 交互式问答 # ==================================================================== def interactive_query(rag_system: OptimizedRAGSystem): """交互式问答模式""" print("\n" + "=" * 60) print("RAG 全链路优化系统 - 交互式问答") print("=" * 60) print("检索模式: none / rewrite / decompose / hyde / multi_query") print("命令: 'mode <模式名>' 切换模式, 'compare' 对比模式, 'quit' 退出") print() current_mode = "rewrite" while True: try: user_input = input(f"[{current_mode}] 你的问题: ").strip() if user_input.lower() in ['quit', 'exit', 'q']: print("\n再见!") break if not user_input: continue # 切换模式命令 if user_input.startswith("mode "): new_mode = user_input[5:].strip() valid_modes = ["none", "rewrite", "decompose", "hyde", "multi_query"] if new_mode in valid_modes: current_mode = new_mode print(f"已切换到模式: {current_mode}") else: print(f"无效模式。可选: {valid_modes}") continue # 对比命令 if user_input == "compare": question = input("输入对比问题: ").strip() if question: rag_system.compare_retrieval_modes(question) continue # 正常查询 print("\n检索并生成中...\n") result = rag_system.query(user_input, mode=current_mode) print("-" * 60) print(f"回答:\n{result['answer']}") print(f"\n[召回文档数: {result['doc_count']}]") print("-" * 60) print() except KeyboardInterrupt: print("\n\n再见!") break except Exception as e: log_error(f"查询出错: {e}") # ==================================================================== # 主函数 # ==================================================================== def main(): """主函数""" # PDF 文件路径 pdf_path = r"./car_info.pdf" # Milvus 配置 milvus_uri = "http://localhost:19530" collection_name = "car_info_optimized" # 检查 API Key 配置 openai_api_key = os.getenv("OPENAI_API_KEY") dashscope_api_key = os.getenv("DASHSCOPE_API_KEY") if not openai_api_key: print("错误: 未设置 OPENAI_API_KEY 环境变量") print("请在 .env 文件中配置: OPENAI_API_KEY=your-api-key") return if not dashscope_api_key: print("错误: 未设置 DASHSCOPE_API_KEY 环境变量") print("请在 .env 文件中配置: DASHSCOPE_API_KEY=your-api-key") return # 检查 PDF 文件 if not os.path.exists(pdf_path): print(f"错误: PDF 文件不存在: {pdf_path}") return try: # 创建优化版 RAG 系统 rag = OptimizedRAGSystem( pdf_path=pdf_path, collection_name=collection_name, milvus_uri=milvus_uri, # 语义切块需要较多 Embedding API 调用 # 如果 API 调用有限制,可以设为 False 使用递归字符切块 use_semantic_chunking=True, ) # 构建知识库 print("\n是否重新构建知识库?(y/n,默认 n 加载已有数据): ", end="") choice = input().strip().lower() if choice == 'y': rag.build_knowledge_base() else: try: rag.load_vectorstore() # 加载文档用于 BM25 检索器 rag.load_pdf() rag.clean_all_pages() rag.split_documents() log_info("已加载现有向量库和文档") except Exception as e: log_warn(f"加载失败: {e},重新构建...") rag.build_knowledge_base() # 进入交互式问答 interactive_query(rag) except Exception as e: log_error(f"系统运行出错: {e}") import traceback traceback.print_exc() if __name__ == "__main__": main()