02-2.RAG_optimize_task.py 27 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800
  1. """
  2. RAG 全链路优化系统
  3. 优化点:
  4. 1. 语义切块(SemanticChunker)替代固定切片
  5. 2. 多路召回(BM25 + 向量检索 + 混合检索)
  6. 3. 查询侧优化(查询重写、查询分解、HyDE)
  7. 4. Reranker 重排序精排
  8. 5. Milvus 向量数据库替代 Chroma
  9. """
  10. import os
  11. import re
  12. import json
  13. from typing import List, Optional
  14. from concurrent.futures import ThreadPoolExecutor
  15. from dotenv import load_dotenv
  16. # LangChain 核心组件
  17. from langchain_community.document_loaders import PyMuPDFLoader
  18. from langchain_text_splitters import RecursiveCharacterTextSplitter
  19. from langchain_community.embeddings import DashScopeEmbeddings
  20. from langchain_community.vectorstores import Milvus
  21. from langchain_community.retrievers import BM25Retriever
  22. from langchain_classic.retrievers import EnsembleRetriever
  23. from langchain_classic.retrievers.multi_query import MultiQueryRetriever
  24. from langchain_core.prompts import ChatPromptTemplate, PromptTemplate
  25. from langchain_openai import ChatOpenAI
  26. from langchain_core.output_parsers import StrOutputParser
  27. from langchain_core.documents import Document
  28. from langchain_experimental.text_splitter import SemanticChunker
  29. # 加载环境变量
  30. load_dotenv()
  31. # ====================================================================
  32. # 日志工具
  33. # ====================================================================
  34. def log_info(message: str):
  35. """打印日志信息"""
  36. print(f"[INFO] {message}")
  37. def log_error(message: str):
  38. """打印错误日志"""
  39. print(f"[ERROR] {message}")
  40. def log_warn(message: str):
  41. """打印警告日志"""
  42. print(f"[WARN] {message}")
  43. # ====================================================================
  44. # 查询侧优化模块
  45. # ====================================================================
  46. class QueryOptimizer:
  47. """查询侧优化器:查询重写、查询分解、HyDE"""
  48. def __init__(self, llm: ChatOpenAI):
  49. self.llm = llm
  50. self._init_chains()
  51. def _init_chains(self):
  52. """初始化各优化链"""
  53. # 查询重写 Prompt
  54. self.rewrite_prompt = PromptTemplate(
  55. input_variables=["query"],
  56. template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。
  57. 改写要求:
  58. 1. 补充隐含的上下文信息
  59. 2. 将口语化表达转为专业表述
  60. 3. 消除歧义,明确查询意图
  61. 4. 保持原意不变,不要添加原问题未提及的内容
  62. 5. 直接输出改写后的查询,不要解释
  63. 用户问题: {query}
  64. 改写后的查询:"""
  65. )
  66. self.rewrite_chain = self.rewrite_prompt | self.llm | StrOutputParser()
  67. # 查询分解 Prompt
  68. self.decompose_prompt = PromptTemplate(
  69. input_variables=["question"],
  70. template="""你是一个问题分解助手。请将用户的复杂问题分解为 2-4 个独立的子问题,
  71. 每个子问题应该能独立检索和回答。
  72. 要求:
  73. 1. 子问题之间互不依赖,可以并行检索
  74. 2. 子问题覆盖原始问题的所有方面
  75. 3. 每个子问题简洁明确
  76. 4. 以 JSON 数组格式输出
  77. 用户问题: {question}
  78. 输出格式: ["子问题1", "子问题2", "子问题3"]
  79. 子问题列表:"""
  80. )
  81. self.decompose_chain = self.decompose_prompt | self.llm | StrOutputParser()
  82. # HyDE Prompt(假设性文档嵌入)
  83. self.hyde_prompt = PromptTemplate(
  84. input_variables=["question"],
  85. template="""请根据以下问题,写一段可能出现在相关文档中的内容(约100-200字)。
  86. 不需要完全准确,只需要包含相关的关键词和表述方式。
  87. 问题: {question}
  88. 假设性文档内容:"""
  89. )
  90. self.hyde_chain = self.hyde_prompt | self.llm | StrOutputParser()
  91. def rewrite(self, query: str) -> str:
  92. """查询重写:将口语化问题转为精确检索查询"""
  93. try:
  94. rewritten = self.rewrite_chain.invoke({"query": query})
  95. log_info(f"查询重写: '{query}' → '{rewritten}'")
  96. return rewritten
  97. except Exception as e:
  98. log_warn(f"查询重写失败: {e},使用原始查询")
  99. return query
  100. def decompose(self, question: str) -> List[str]:
  101. """查询分解:将复杂问题拆分为子问题"""
  102. try:
  103. result = self.decompose_chain.invoke({"question": question})
  104. sub_queries = json.loads(result.strip())
  105. log_info(f"查询分解: '{question}' → {len(sub_queries)} 个子问题")
  106. return sub_queries
  107. except (json.JSONDecodeError, Exception) as e:
  108. log_warn(f"查询分解失败: {e},使用原始问题")
  109. return [question]
  110. def hyde(self, question: str) -> str:
  111. """HyDE:生成假设性文档用于检索"""
  112. try:
  113. hypothetical_doc = self.hyde_chain.invoke({"question": question})
  114. log_info(f"HyDE 生成假设文档: {hypothetical_doc[:80]}...")
  115. return hypothetical_doc
  116. except Exception as e:
  117. log_warn(f"HyDE 生成失败: {e},使用原始问题")
  118. return question
  119. # ====================================================================
  120. # Reranker 重排序模块
  121. # ====================================================================
  122. class LLMReranker:
  123. """基于 LLM 的重排序器(无需额外模型,用大模型做相关性打分)"""
  124. def __init__(self, llm: ChatOpenAI):
  125. self.llm = llm
  126. self.score_prompt = PromptTemplate(
  127. input_variables=["query", "document"],
  128. template="""请评估以下文档与用户问题的相关性,给出 0-10 的分数。
  129. 只输出数字分数,不要任何解释。
  130. 用户问题: {query}
  131. 文档内容: {document}
  132. 相关性分数(0-10):"""
  133. )
  134. self.score_chain = self.score_prompt | self.llm | StrOutputParser()
  135. def rerank(self, query: str, documents: List[Document], top_k: int = 3) -> List[Document]:
  136. """对检索结果重排序,返回最相关的 top_k 个文档"""
  137. if not documents:
  138. return []
  139. if len(documents) <= top_k:
  140. return documents
  141. log_info(f"Reranker: 对 {len(documents)} 个文档重排序,保留 Top-{top_k}")
  142. scored_docs = []
  143. for doc in documents:
  144. try:
  145. score_str = self.score_chain.invoke({
  146. "query": query,
  147. "document": doc.page_content[:500]
  148. })
  149. # 提取数字分数
  150. score = float(re.search(r'[\d.]+', score_str).group())
  151. scored_docs.append((doc, score))
  152. except Exception:
  153. # 打分失败时给默认分数
  154. scored_docs.append((doc, 5.0))
  155. # 按分数降序排序
  156. scored_docs.sort(key=lambda x: x[1], reverse=True)
  157. # 返回 top_k 个
  158. result = [doc for doc, score in scored_docs[:top_k]]
  159. log_info(f"Reranker 完成,分数范围: {scored_docs[0][1]:.1f} - {scored_docs[-1][1]:.1f}")
  160. return result
  161. # ====================================================================
  162. # 核心 RAG 系统(全链路优化版)
  163. # ====================================================================
  164. class OptimizedRAGSystem:
  165. """
  166. 全链路优化 RAG 系统
  167. 优化链路:
  168. 文档加载 → 数据清洗 → 语义切块 → Milvus存储 → 多路召回 → Reranker精排 → 生成回答
  169. 查询优化:
  170. 用户问题 → 查询重写/分解/HyDE → 多路召回 → 重排序 → 生成
  171. """
  172. def __init__(
  173. self,
  174. pdf_path: str,
  175. collection_name: str = "rag_optimized",
  176. milvus_uri: str = "http://localhost:19530",
  177. use_semantic_chunking: bool = True,
  178. ):
  179. """
  180. 初始化优化版 RAG 系统
  181. Args:
  182. pdf_path: PDF 文件路径
  183. collection_name: Milvus 集合名称
  184. milvus_uri: Milvus 连接地址
  185. use_semantic_chunking: 是否使用语义切块(需要较多 API 调用)
  186. """
  187. self.pdf_path = pdf_path
  188. self.collection_name = collection_name
  189. self.milvus_uri = milvus_uri
  190. self.use_semantic_chunking = use_semantic_chunking
  191. self.pages = None
  192. self.split_docs = None
  193. self.vectorstore = None
  194. # 初始化 Embedding 模型
  195. self.embedding_model = DashScopeEmbeddings(
  196. model=os.getenv("EMBEDDING_MODEL", "text-embedding-v3"),
  197. dashscope_api_key=os.getenv("DASHSCOPE_API_KEY", "")
  198. )
  199. # 初始化大模型
  200. self.llm = ChatOpenAI(
  201. model=os.getenv("MODEL_NAME", "deepseek-chat"),
  202. openai_api_key=os.getenv("OPENAI_API_KEY", ""),
  203. openai_api_base=os.getenv("OPENAI_API_BASE", "https://api.deepseek.com/v1"),
  204. temperature=0.7
  205. )
  206. # 初始化查询优化器和重排序器
  207. self.query_optimizer = QueryOptimizer(self.llm)
  208. self.reranker = LLMReranker(self.llm)
  209. log_info("全链路优化 RAG 系统初始化完成")
  210. log_info(f"模型: {os.getenv('MODEL_NAME', 'deepseek-chat')}")
  211. log_info(f"Milvus: {milvus_uri}")
  212. log_info(f"切块策略: {'语义切块' if use_semantic_chunking else '递归字符切块'}")
  213. # ================================================================
  214. # 文档处理阶段
  215. # ================================================================
  216. def load_pdf(self) -> List[Document]:
  217. """加载 PDF 文档"""
  218. log_info(f"加载 PDF: {self.pdf_path}")
  219. if not os.path.exists(self.pdf_path):
  220. raise FileNotFoundError(f"PDF 文件不存在: {self.pdf_path}")
  221. loader = PyMuPDFLoader(self.pdf_path)
  222. self.pages = loader.load()
  223. total_content = sum(len(p.page_content) for p in self.pages)
  224. if total_content == 0:
  225. raise ValueError("PDF 文件无可提取文本内容,请使用文本型 PDF 或进行 OCR 处理")
  226. log_info(f"加载完成: {len(self.pages)} 页, {total_content} 字符")
  227. return self.pages
  228. def clean_text(self, text: str) -> str:
  229. """清洗文本:去除噪声、规范化"""
  230. if not text or not text.strip():
  231. return ""
  232. # 删除多余换行(保留最多2个)
  233. text = re.sub(r'\n{3,}', '\n\n', text)
  234. # 删除项目符号
  235. text = text.replace('•', '').replace('·', '')
  236. # 合并多余空格
  237. text = re.sub(r'[^\S\n]+', ' ', text)
  238. # 删除控制字符
  239. text = re.sub(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f-\x9f]', '', text)
  240. return text.strip()
  241. def clean_all_pages(self) -> List[Document]:
  242. """清洗所有页面"""
  243. if not self.pages:
  244. raise ValueError("请先调用 load_pdf()")
  245. log_info("清洗文档...")
  246. for page in self.pages:
  247. page.page_content = self.clean_text(page.page_content)
  248. # 过滤空页
  249. self.pages = [p for p in self.pages if p.page_content.strip()]
  250. log_info(f"清洗完成: {len(self.pages)} 页有效内容")
  251. return self.pages
  252. def split_documents(self) -> List[Document]:
  253. """
  254. 文档切块(支持语义切块和递归字符切块两种策略)
  255. 语义切块:基于相邻句子的语义相似度,在话题转变处切分
  256. 递归字符切块:按字符数切分,带重叠防止上下文断裂
  257. """
  258. if not self.pages:
  259. raise ValueError("请先调用 load_pdf()")
  260. if self.use_semantic_chunking:
  261. log_info("使用语义切块策略(SemanticChunker)...")
  262. try:
  263. text_splitter = SemanticChunker(
  264. embeddings=self.embedding_model,
  265. breakpoint_threshold_type="percentile",
  266. # 90 表示只有相似度排名后 10% 的位置才切分
  267. # 值越大切得越少,块越大
  268. breakpoint_threshold_amount=90,
  269. )
  270. self.split_docs = text_splitter.split_documents(self.pages)
  271. except Exception as e:
  272. log_warn(f"语义切块失败: {e},回退到递归字符切块")
  273. self.split_docs = self._fallback_split()
  274. else:
  275. self.split_docs = self._fallback_split()
  276. # 过滤空块
  277. self.split_docs = [
  278. doc for doc in self.split_docs
  279. if doc.page_content and doc.page_content.strip()
  280. ]
  281. log_info(f"切块完成: {len(self.split_docs)} 个文档块")
  282. if self.split_docs:
  283. avg_len = sum(len(d.page_content) for d in self.split_docs) / len(self.split_docs)
  284. log_info(f"平均块长度: {avg_len:.0f} 字符")
  285. return self.split_docs
  286. def _fallback_split(self) -> List[Document]:
  287. """回退策略:递归字符切块"""
  288. log_info("使用递归字符切块策略...")
  289. text_splitter = RecursiveCharacterTextSplitter(
  290. separators=["\n\n", "\n", "。", "!", "?", ";", " ", ""],
  291. chunk_size=300,
  292. chunk_overlap=50,
  293. length_function=len,
  294. )
  295. return text_splitter.split_documents(self.pages)
  296. # ================================================================
  297. # 向量存储与多路召回
  298. # ================================================================
  299. def create_vectorstore(self) -> Milvus:
  300. """将文档块存入 Milvus 向量数据库"""
  301. if not self.split_docs:
  302. raise ValueError("请先调用 split_documents()")
  303. log_info(f"创建 Milvus 向量库: {self.collection_name}")
  304. self.vectorstore = Milvus.from_documents(
  305. documents=self.split_docs,
  306. embedding=self.embedding_model,
  307. connection_args={"uri": self.milvus_uri},
  308. collection_name=self.collection_name,
  309. drop_old=True, # 开发阶段覆盖旧数据
  310. )
  311. log_info(f"Milvus 集合创建完成: {len(self.split_docs)} 个文档块已入库")
  312. return self.vectorstore
  313. def load_vectorstore(self) -> Milvus:
  314. """加载已有的 Milvus 向量数据库"""
  315. log_info(f"加载 Milvus 集合: {self.collection_name}")
  316. self.vectorstore = Milvus(
  317. embedding_function=self.embedding_model,
  318. connection_args={"uri": self.milvus_uri},
  319. collection_name=self.collection_name,
  320. )
  321. log_info("Milvus 向量库加载完成")
  322. return self.vectorstore
  323. def build_ensemble_retriever(self, k: int = 10) -> EnsembleRetriever:
  324. """
  325. 构建混合检索器(BM25 + 向量检索)
  326. BM25: 关键词精确匹配,擅长处理精确术语
  327. 向量检索: 语义相似度匹配,擅长处理同义词和语义理解
  328. 混合: 两者加权融合,取长补短
  329. """
  330. if not self.vectorstore or not self.split_docs:
  331. raise ValueError("请先创建向量库")
  332. # BM25 检索器(关键词匹配)
  333. bm25_retriever = BM25Retriever.from_documents(self.split_docs)
  334. bm25_retriever.k = k
  335. # 向量检索器(语义匹配)
  336. vector_retriever = self.vectorstore.as_retriever(
  337. search_kwargs={"k": k}
  338. )
  339. # 混合检索器:BM25 权重 0.4,向量检索权重 0.6
  340. ensemble_retriever = EnsembleRetriever(
  341. retrievers=[bm25_retriever, vector_retriever],
  342. weights=[0.4, 0.6],
  343. )
  344. log_info(f"混合检索器构建完成 (BM25:0.4 + Vector:0.6, Top-{k})")
  345. return ensemble_retriever
  346. def build_multi_query_retriever(
  347. self, base_retriever: EnsembleRetriever
  348. ) -> MultiQueryRetriever:
  349. """
  350. 构建多查询检索器
  351. 原理:用 LLM 将一个问题改写为多个不同角度的问题,
  352. 分别检索后合并去重,增加召回率
  353. """
  354. multi_query_prompt = PromptTemplate(
  355. input_variables=["question"],
  356. template="""你是一个问题改写助手。请为下面的问题生成 3 个不同的改写版本。
  357. 每个版本应从不同角度表达相同意图,使用不同的关键词。
  358. 每个问题单独一行,不要编号。
  359. 原始问题: {question}
  360. 改写后的问题:"""
  361. )
  362. retriever = MultiQueryRetriever.from_llm(
  363. retriever=base_retriever,
  364. llm=self.llm,
  365. prompt=multi_query_prompt,
  366. )
  367. log_info("多查询检索器构建完成")
  368. return retriever
  369. # ================================================================
  370. # 检索与生成
  371. # ================================================================
  372. def retrieve_with_optimization(
  373. self,
  374. query: str,
  375. retriever,
  376. mode: str = "rewrite",
  377. rerank_top_k: int = 3,
  378. ) -> List[Document]:
  379. """
  380. 带查询优化 + Reranker 的检索流程
  381. Args:
  382. query: 用户原始问题
  383. retriever: 检索器实例
  384. mode: 查询优化模式
  385. - "none": 不做优化,直接检索
  386. - "rewrite": 查询重写后检索
  387. - "decompose": 查询分解后并行检索
  388. - "hyde": HyDE 假设文档检索
  389. - "multi_query": 多查询检索(需要传入 MultiQueryRetriever)
  390. rerank_top_k: Reranker 精排后保留的文档数
  391. Returns:
  392. 精排后的文档列表
  393. """
  394. log_info(f"检索模式: {mode} | 问题: {query}")
  395. # 第一步:查询优化
  396. if mode == "rewrite":
  397. optimized_query = self.query_optimizer.rewrite(query)
  398. raw_docs = retriever.invoke(optimized_query)
  399. elif mode == "decompose":
  400. sub_queries = self.query_optimizer.decompose(query)
  401. raw_docs = self._parallel_retrieve(sub_queries, retriever)
  402. elif mode == "hyde":
  403. hypothetical_doc = self.query_optimizer.hyde(query)
  404. raw_docs = retriever.invoke(hypothetical_doc)
  405. elif mode == "multi_query":
  406. # MultiQueryRetriever 内部自带多查询逻辑
  407. raw_docs = retriever.invoke(query)
  408. else:
  409. # mode == "none"
  410. raw_docs = retriever.invoke(query)
  411. log_info(f"初步召回: {len(raw_docs)} 个文档")
  412. # 第二步:Reranker 精排
  413. if len(raw_docs) > rerank_top_k:
  414. reranked_docs = self.reranker.rerank(query, raw_docs, top_k=rerank_top_k)
  415. else:
  416. reranked_docs = raw_docs
  417. log_info(f"精排后: {len(reranked_docs)} 个文档")
  418. return reranked_docs
  419. def _parallel_retrieve(
  420. self, queries: List[str], retriever
  421. ) -> List[Document]:
  422. """并行检索多个查询并合并去重"""
  423. all_docs = []
  424. seen_contents = set()
  425. def retrieve_one(q):
  426. return retriever.invoke(q)
  427. with ThreadPoolExecutor(max_workers=min(len(queries), 4)) as executor:
  428. futures = [executor.submit(retrieve_one, q) for q in queries]
  429. for future in futures:
  430. try:
  431. docs = future.result()
  432. for doc in docs:
  433. content_hash = hash(doc.page_content)
  434. if content_hash not in seen_contents:
  435. seen_contents.add(content_hash)
  436. all_docs.append(doc)
  437. except Exception as e:
  438. log_warn(f"并行检索子任务失败: {e}")
  439. return all_docs
  440. def generate_answer(self, query: str, relevant_docs: List[Document]) -> str:
  441. """基于检索结果生成回答"""
  442. if not relevant_docs:
  443. return "抱歉,根据现有资料,我找不到与您问题相关的信息。"
  444. log_info("生成回答...")
  445. context = "\n\n---\n\n".join([doc.page_content for doc in relevant_docs])
  446. prompt = ChatPromptTemplate.from_template("""你是一个专业的知识库助手。请根据以下检索到的上下文回答用户问题。
  447. **规则:**
  448. - 只基于提供的上下文回答,不要编造信息
  449. - 如果上下文中没有相关信息,请明确说明
  450. - 回答要简洁直接,条理清晰
  451. - 引用原文时用引号标注
  452. **检索到的上下文:**
  453. {context}
  454. **用户问题:**
  455. {question}
  456. **回答:**""")
  457. chain = prompt | self.llm | StrOutputParser()
  458. answer = chain.invoke({"context": context, "question": query})
  459. log_info("回答生成完成")
  460. return answer
  461. def query(
  462. self,
  463. question: str,
  464. mode: str = "rewrite",
  465. rerank_top_k: int = 3,
  466. ) -> dict:
  467. """
  468. 完整查询流程:查询优化 → 多路召回 → 重排序 → 生成回答
  469. Args:
  470. question: 用户问题
  471. mode: 查询优化模式 ("none"/"rewrite"/"decompose"/"hyde"/"multi_query")
  472. rerank_top_k: 精排后保留文档数
  473. Returns:
  474. 包含回答和中间结果的字典
  475. """
  476. # 构建检索器
  477. ensemble_retriever = self.build_ensemble_retriever(k=10)
  478. # 如果是 multi_query 模式,额外包装一层
  479. if mode == "multi_query":
  480. retriever = self.build_multi_query_retriever(ensemble_retriever)
  481. else:
  482. retriever = ensemble_retriever
  483. # 检索 + 重排序
  484. relevant_docs = self.retrieve_with_optimization(
  485. query=question,
  486. retriever=retriever,
  487. mode=mode,
  488. rerank_top_k=rerank_top_k,
  489. )
  490. # 生成回答
  491. answer = self.generate_answer(question, relevant_docs)
  492. return {
  493. "question": question,
  494. "mode": mode,
  495. "doc_count": len(relevant_docs),
  496. "answer": answer,
  497. "sources": [doc.page_content[:100] for doc in relevant_docs],
  498. }
  499. # ================================================================
  500. # 构建与对比
  501. # ================================================================
  502. def build_knowledge_base(self):
  503. """构建完整知识库:加载 → 清洗 → 切块 → 向量化"""
  504. log_info("=" * 60)
  505. log_info("开始构建优化版知识库...")
  506. log_info("=" * 60)
  507. self.load_pdf()
  508. self.clean_all_pages()
  509. self.split_documents()
  510. self.create_vectorstore()
  511. log_info("=" * 60)
  512. log_info("知识库构建完成!")
  513. log_info("=" * 60)
  514. def compare_retrieval_modes(self, question: str):
  515. """对比不同检索模式的效果"""
  516. log_info("=" * 60)
  517. log_info(f"对比测试: {question}")
  518. log_info("=" * 60)
  519. modes = ["none", "rewrite", "decompose", "hyde"]
  520. results = {}
  521. for mode in modes:
  522. log_info(f"\n--- 模式: {mode} ---")
  523. try:
  524. result = self.query(question, mode=mode)
  525. results[mode] = result
  526. print(f"\n[{mode}] 召回 {result['doc_count']} 个文档")
  527. print(f"回答: {result['answer'][:200]}...")
  528. except Exception as e:
  529. log_error(f"模式 {mode} 执行失败: {e}")
  530. results[mode] = {"error": str(e)}
  531. return results
  532. # ====================================================================
  533. # 交互式问答
  534. # ====================================================================
  535. def interactive_query(rag_system: OptimizedRAGSystem):
  536. """交互式问答模式"""
  537. print("\n" + "=" * 60)
  538. print("RAG 全链路优化系统 - 交互式问答")
  539. print("=" * 60)
  540. print("检索模式: none / rewrite / decompose / hyde / multi_query")
  541. print("命令: 'mode <模式名>' 切换模式, 'compare' 对比模式, 'quit' 退出")
  542. print()
  543. current_mode = "rewrite"
  544. while True:
  545. try:
  546. user_input = input(f"[{current_mode}] 你的问题: ").strip()
  547. if user_input.lower() in ['quit', 'exit', 'q']:
  548. print("\n再见!")
  549. break
  550. if not user_input:
  551. continue
  552. # 切换模式命令
  553. if user_input.startswith("mode "):
  554. new_mode = user_input[5:].strip()
  555. valid_modes = ["none", "rewrite", "decompose", "hyde", "multi_query"]
  556. if new_mode in valid_modes:
  557. current_mode = new_mode
  558. print(f"已切换到模式: {current_mode}")
  559. else:
  560. print(f"无效模式。可选: {valid_modes}")
  561. continue
  562. # 对比命令
  563. if user_input == "compare":
  564. question = input("输入对比问题: ").strip()
  565. if question:
  566. rag_system.compare_retrieval_modes(question)
  567. continue
  568. # 正常查询
  569. print("\n检索并生成中...\n")
  570. result = rag_system.query(user_input, mode=current_mode)
  571. print("-" * 60)
  572. print(f"回答:\n{result['answer']}")
  573. print(f"\n[召回文档数: {result['doc_count']}]")
  574. print("-" * 60)
  575. print()
  576. except KeyboardInterrupt:
  577. print("\n\n再见!")
  578. break
  579. except Exception as e:
  580. log_error(f"查询出错: {e}")
  581. # ====================================================================
  582. # 主函数
  583. # ====================================================================
  584. def main():
  585. """主函数"""
  586. # PDF 文件路径
  587. pdf_path = r"./car_info.pdf"
  588. # Milvus 配置
  589. milvus_uri = "http://localhost:19530"
  590. collection_name = "car_info_optimized"
  591. # 检查 API Key 配置
  592. openai_api_key = os.getenv("OPENAI_API_KEY")
  593. dashscope_api_key = os.getenv("DASHSCOPE_API_KEY")
  594. if not openai_api_key:
  595. print("错误: 未设置 OPENAI_API_KEY 环境变量")
  596. print("请在 .env 文件中配置: OPENAI_API_KEY=your-api-key")
  597. return
  598. if not dashscope_api_key:
  599. print("错误: 未设置 DASHSCOPE_API_KEY 环境变量")
  600. print("请在 .env 文件中配置: DASHSCOPE_API_KEY=your-api-key")
  601. return
  602. # 检查 PDF 文件
  603. if not os.path.exists(pdf_path):
  604. print(f"错误: PDF 文件不存在: {pdf_path}")
  605. return
  606. try:
  607. # 创建优化版 RAG 系统
  608. rag = OptimizedRAGSystem(
  609. pdf_path=pdf_path,
  610. collection_name=collection_name,
  611. milvus_uri=milvus_uri,
  612. # 语义切块需要较多 Embedding API 调用
  613. # 如果 API 调用有限制,可以设为 False 使用递归字符切块
  614. use_semantic_chunking=True,
  615. )
  616. # 构建知识库
  617. print("\n是否重新构建知识库?(y/n,默认 n 加载已有数据): ", end="")
  618. choice = input().strip().lower()
  619. if choice == 'y':
  620. rag.build_knowledge_base()
  621. else:
  622. try:
  623. rag.load_vectorstore()
  624. # 加载文档用于 BM25 检索器
  625. rag.load_pdf()
  626. rag.clean_all_pages()
  627. rag.split_documents()
  628. log_info("已加载现有向量库和文档")
  629. except Exception as e:
  630. log_warn(f"加载失败: {e},重新构建...")
  631. rag.build_knowledge_base()
  632. # 进入交互式问答
  633. interactive_query(rag)
  634. except Exception as e:
  635. log_error(f"系统运行出错: {e}")
  636. import traceback
  637. traceback.print_exc()
  638. if __name__ == "__main__":
  639. main()