rag_agent_optimized-two.py 10 KB

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