rag_agent_optimized.py 9.6 KB

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