rag_agent_optimized-hyde-search.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425
  1. import os
  2. import chromadb
  3. from dotenv import load_dotenv
  4. from langchain_community.document_loaders import PyMuPDFLoader
  5. from langchain_community.embeddings import DashScopeEmbeddings
  6. from langchain_community.retrievers import BM25Retriever
  7. from schema import RAGWithQueryRewriting
  8. from schema import RAGWithDecomposition
  9. from langchain_community.vectorstores import Chroma
  10. from langchain_core.output_parsers import StrOutputParser
  11. from langchain_core.prompts import ChatPromptTemplate
  12. from langchain_openai import ChatOpenAI
  13. from langchain_community.retrievers import BM25Retriever
  14. from langchain_community.embeddings import DashScopeEmbeddings
  15. from langchain_classic.retrievers import EnsembleRetriever
  16. from langchain_classic.chains.combine_documents import create_stuff_documents_chain
  17. from langchain_text_splitters import RecursiveCharacterTextSplitter
  18. from langchain_core.prompts import PromptTemplate
  19. from langchain_community.chat_models import ChatTongyi
  20. from langchain_classic.retrievers.multi_query import MultiQueryRetriever
  21. from langchain_experimental.text_splitter import SemanticChunker
  22. # ── 配置 ──────────────────────────────────────────────────────────
  23. BASE_DIR = os.path.dirname(os.path.abspath(__file__))
  24. PDF_PATH = os.path.join(BASE_DIR, "docs", "浦发上海浦东发展银行西安分行个金客户经理考核办法.pdf")
  25. CHROMA_DB_PATH = os.path.join(BASE_DIR, "chroma")
  26. COLLECTION_NAME = "shanghai_bank_policy"
  27. PROMPT_TEMPLATE = """
  28. 你是一个专业的知识库助手。请根据以下上下文回答问题。
  29. **规则:**
  30. - 只基于提供的上下文回答,不要编造
  31. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  32. - 回答要简洁直接,引用原文时用引号
  33. **上下文:**
  34. {context}
  35. **问题:**
  36. {question}
  37. """
  38. def load_config() -> dict:
  39. """加载 .env 中的配置项"""
  40. load_dotenv()
  41. config = {
  42. "api_key": os.getenv("ALIYUN_API_KEY"),
  43. "base_url": os.getenv("ALIYUN_BASE_URL"),
  44. "chat_model": os.getenv("ALIYUN_CHAT_MODEL"),
  45. "embedding_model": os.getenv("ALIYUN_EMBEDDING_MODEL", "text-embedding-v3"),
  46. }
  47. missing = [k for k, v in config.items() if not v and k != "embedding_model"]
  48. if missing:
  49. raise EnvironmentError(f"缺少环境变量: {missing},请检查 .env 文件")
  50. return config
  51. def init_llm(config: dict) -> ChatOpenAI:
  52. return ChatOpenAI(
  53. base_url=config["base_url"],
  54. api_key=config["api_key"],
  55. model=config["chat_model"],
  56. temperature=0,
  57. timeout=120,
  58. max_retries=1,
  59. )
  60. def init_embeddings(config: dict) -> DashScopeEmbeddings:
  61. return DashScopeEmbeddings(
  62. model=config["embedding_model"],
  63. dashscope_api_key=config["api_key"],
  64. )
  65. def clean_pdf_text(text: str) -> str:
  66. """清洗 PDF 解析出的文本,去除常见噪声"""
  67. import re
  68. # 删除非中文字符之间的换行符
  69. text = re.sub(r'[^一](\n)[^一]',
  70. lambda m: m.group(0).replace('\n', ''), text)
  71. # 删除项目符号和多余空格
  72. text = text.replace('•', '').replace(' ', ' ')
  73. # 删除连续的换行符(保留一个)
  74. text = re.sub(r'\n{2,}', '\n', text)
  75. return text.strip()
  76. def build_or_load_vectorstore(embeddings: DashScopeEmbeddings) -> Chroma:
  77. """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
  78. chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
  79. existing_collections = [col.name for col in chroma_client.list_collections()]
  80. if COLLECTION_NAME in existing_collections:
  81. print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
  82. return Chroma(
  83. embedding_function=embeddings,
  84. collection_name=COLLECTION_NAME,
  85. client=chroma_client,
  86. )
  87. print(f"ℹ️ 未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
  88. loader = PyMuPDFLoader(PDF_PATH)
  89. pages = loader.load()
  90. # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
  91. pages = clean_documents(pages)
  92. # 创建语义分块器
  93. # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
  94. # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
  95. text_splitter = SemanticChunker(
  96. embeddings=embeddings,
  97. breakpoint_threshold_type="percentile",
  98. breakpoint_threshold_amount=95 # 值越大,切出来的块越少(越粗)
  99. )
  100. docs = text_splitter.split_documents(pages)
  101. vectorstore = Chroma.from_documents(
  102. embedding=embeddings,
  103. collection_name=COLLECTION_NAME,
  104. client=chroma_client,
  105. documents=docs,
  106. )
  107. print(f"✅ 建库完成,共 {len(docs)} 个分块")
  108. return vectorstore
  109. def build_ensemble_retriever(embeddings: DashScopeEmbeddings) -> EnsembleRetriever:
  110. """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
  111. chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
  112. existing_collections = [col.name for col in chroma_client.list_collections()]
  113. loader = PyMuPDFLoader(PDF_PATH)
  114. pages = loader.load()
  115. # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
  116. pages = clean_documents(pages)
  117. # 创建语义分块器
  118. # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
  119. # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
  120. text_splitter = SemanticChunker(
  121. embeddings=embeddings,
  122. breakpoint_threshold_type="percentile",
  123. breakpoint_threshold_amount=95 # 值越大,切出来的块越少(越粗)
  124. )
  125. docs = text_splitter.split_documents(pages)
  126. vectorstore = None
  127. if COLLECTION_NAME in existing_collections:
  128. print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
  129. vectorstore = Chroma(
  130. embedding_function=embeddings,
  131. collection_name=COLLECTION_NAME,
  132. client=chroma_client,
  133. )
  134. else:
  135. print(f"ℹ️ 未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
  136. vectorstore = Chroma.from_documents(
  137. embedding=embeddings,
  138. collection_name=COLLECTION_NAME,
  139. client=chroma_client,
  140. documents=docs,
  141. )
  142. # ---- 创建 BM25 检索器 ----
  143. # BM25 不需要向量,直接基于文本的关键词匹配
  144. bm25_retriever = BM25Retriever.from_documents(docs)
  145. bm25_retriever.k = 10
  146. vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
  147. ensemble_retriever = EnsembleRetriever(
  148. retrievers=[bm25_retriever, vector_retriever],
  149. weights=[0.4, 0.6],
  150. normalize_scores=True # 将不同检索器的分数归一化到 [0,1],避免偏差
  151. )
  152. return ensemble_retriever
  153. def build_multi_query_retriever(embeddings: DashScopeEmbeddings, llm) -> MultiQueryRetriever:
  154. # 自定义改写提示词(可选,不写则用默认的)
  155. CUSTOM_PROMPT = PromptTemplate(
  156. input_variables=["question"],
  157. template="""你是一个专业的问题改写助手。请为下面的问题生成 4 个不同的改写版本,
  158. 每个版本应该:
  159. - 从不同角度表达相同的意图
  160. - 使用不同的关键词和表达方式
  161. - 保持问题的核心含义
  162. 每个问题单独一行,不要编号。
  163. 原始问题: {question}
  164. 改写后的问题:"""
  165. )
  166. ensemble_retriever = build_ensemble_retriever(embeddings=embeddings)
  167. # 创建多查询检索器
  168. # 底层检索器用的是上面的混合检索器,这样每个改写问题都会走混合检索
  169. multi_query_retriever = MultiQueryRetriever.from_llm(
  170. retriever=ensemble_retriever,
  171. llm=llm,
  172. prompt=CUSTOM_PROMPT
  173. )
  174. return multi_query_retriever
  175. def ask(query: str, retriever, llm) -> str:
  176. """检索 + 生成回答"""
  177. relevant_docs = retriever.invoke(query)
  178. context = "\n\n---\n\n".join(d.page_content for d in relevant_docs)
  179. prompt = ChatPromptTemplate.from_template(PROMPT_TEMPLATE)
  180. chain = prompt | llm | StrOutputParser()
  181. return chain.invoke({"context": context, "question": query})
  182. def build_hyde_chain(llm) -> dict:
  183. # HyDE Prompt:生成假设性文档
  184. hyde_prompt = PromptTemplate(
  185. input_variables=["question"],
  186. template="""请根据以下问题,写一段可能包含答案的文档片段。
  187. 要求:
  188. 1. 像真实文档一样专业、详细
  189. 2. 包含具体的数据、步骤或事实
  190. 3. 长度在 100-200 字之间
  191. 4. 即使你不确定答案,也要根据问题合理推测,写出一段"看起来像真的"文档
  192. 问题: {question}
  193. 假设性文档:"""
  194. )
  195. hyde_chain = hyde_prompt | llm | StrOutputParser()
  196. return hyde_chain
  197. def build_rewrite_chain(llm) -> dict:
  198. # 查询重写 Prompt
  199. rewrite_prompt = PromptTemplate(
  200. input_variables=["query"],
  201. template="""你是一个查询优化助手。请将用户的口语化问题改写为更适合信息检索的精确查询。
  202. 改写要求:
  203. 1. 补充隐含的上下文信息
  204. 2. 将口语化表达转为专业表述
  205. 3. 消除歧义,明确查询意图
  206. 4. 保持原意不变,不要添加原问题未提及的内容
  207. 5. 直接输出改写后的查询,不要解释
  208. 用户问题: {query}
  209. 改写后的查询:"""
  210. )
  211. # 构建重写链
  212. rewrite_chain = rewrite_prompt | llm | StrOutputParser()
  213. return rewrite_chain
  214. def hybrid_hyde_search(llm, embeddings: DashScopeEmbeddings, hyde_chain, question) -> dict:
  215. """如果 collection 已存在则直接加载,否则解析 PDF 并建库"""
  216. chroma_client = chromadb.PersistentClient(path=CHROMA_DB_PATH)
  217. existing_collections = [col.name for col in chroma_client.list_collections()]
  218. loader = PyMuPDFLoader(PDF_PATH)
  219. pages = loader.load()
  220. # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
  221. pages = clean_documents(pages)
  222. # 创建语义分块器
  223. # breakpoint_threshold_type="percentile" 表示用百分位数法确定切分阈值
  224. # breakpoint_threshold_amount=95 表示只有相似度排名后 5% 的位置才会被切开
  225. text_splitter = SemanticChunker(
  226. embeddings=embeddings,
  227. breakpoint_threshold_type="percentile",
  228. breakpoint_threshold_amount=95 # 值越大,切出来的块越少(越粗)
  229. )
  230. docs = text_splitter.split_documents(pages)
  231. vectorstore = None
  232. if COLLECTION_NAME in existing_collections:
  233. print(f"✅ 检测到已有 collection「{COLLECTION_NAME}」,直接加载")
  234. vectorstore = Chroma(
  235. embedding_function=embeddings,
  236. collection_name=COLLECTION_NAME,
  237. client=chroma_client,
  238. )
  239. else:
  240. print(f"ℹ️ 未检测到 collection「{COLLECTION_NAME}」,开始解析文档并建库")
  241. vectorstore = Chroma.from_documents(
  242. embedding=embeddings,
  243. collection_name=COLLECTION_NAME,
  244. client=chroma_client,
  245. documents=docs,
  246. )
  247. # ---- 创建 BM25 检索器 ----
  248. # BM25 不需要向量,直接基于文本的关键词匹配
  249. bm25_retriever = BM25Retriever.from_documents(docs)
  250. bm25_retriever.k = 10
  251. hypothetical_doc = hyde_chain.invoke({"question": question})
  252. # 方案一:用假设性文档检索(偏语义)
  253. hyde_docs = vectorstore.similarity_search(hypothetical_doc, k=10)
  254. bm25_docs = bm25_retriever.invoke(hypothetical_doc, k=10) # question must be a str here
  255. # 方案二:用原始问题检索(偏精确)
  256. original_docs = vectorstore.similarity_search(question, k=10)
  257. # 方案二:用原始问题检索(偏精确)
  258. # 合并去重
  259. seen = set()
  260. merged_docs = []
  261. for doc in hyde_docs + bm25_docs:
  262. if hash(doc.page_content) not in seen:
  263. seen.add(hash(doc.page_content))
  264. merged_docs.append(doc)
  265. all_docs = merged_docs[:10]
  266. # 第三步:用原始问题 + 检索结果生成答案
  267. context = "\n\n".join([doc.page_content for doc in all_docs])
  268. answer_prompt = f"""基于以下上下文回答用户问题。如果上下文中没有相关信息,请说明。
  269. 上下文:{context}
  270. 用户问题:{question}
  271. 答案:"""
  272. answer = llm.invoke(answer_prompt).content
  273. return {
  274. "original_query": question,
  275. "hypothetical_doc": hypothetical_doc,
  276. "retrieved_docs": docs,
  277. "answer": answer
  278. }
  279. def clean_pdf_text(text: str) -> str:
  280. """清洗 PDF 解析出的文本,去除常见噪声"""
  281. import re
  282. # 删除非中文字符之间的换行符
  283. text = re.sub(r'[^一](\n)[^一]',
  284. lambda m: m.group(0).replace('\n', ''), text)
  285. # 删除项目符号和多余空格
  286. text = text.replace('•', '').replace(' ', ' ')
  287. # 删除连续的换行符(保留一个)
  288. text = re.sub(r'\n{2,}', '\n', text)
  289. return text.strip()
  290. def clean_documents(pages: list) -> list:
  291. """对 Document 列表逐个清洗 page_content,返回清洗后的新 Document 列表"""
  292. for page in pages:
  293. page.page_content = clean_pdf_text(page.page_content)
  294. return pages
  295. def main():
  296. config = load_config()
  297. llm = init_llm(config)
  298. embeddings = init_embeddings(config)
  299. print("✅ 模型初始化完成")
  300. query = "什么是工作质量考核标准,还有聘任考核程序是什么?"
  301. hyde_chain = build_hyde_chain(llm)
  302. result = hybrid_hyde_search(llm, embeddings, hyde_chain, query)
  303. print("\n" + "=" * 60)
  304. print(f"问题:{query}")
  305. print("=" * 60)
  306. print(result["answer"])
  307. if __name__ == "__main__":
  308. main()