|
@@ -0,0 +1,800 @@
|
|
|
|
|
+"""
|
|
|
|
|
+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()
|