rag_practice.py 8.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205
  1. """
  2. 一. 构建知识库
  3. 1. 收集资料(无优化空间,但为质量的前提)
  4. 2. 解析并清洗数据
  5. ①清洗噪音
  6. ②清洗多余无意义字符
  7. 3. 分块(决定7成左右RAG质量)
  8. ①按固定字符切分
  9. ②按递归符号切分(递归符号:需人工分析,少量可官网AI总结)
  10. ③按文档块标识符切分(一般固定文档格式类型,如markdown)
  11. ④按语义模块切分(费钱,使用embadding模型)
  12. ⑤按LLM切分(烧钱,不考虑钱直接冲)
  13. 4. 向量化(智普embadding模型, 余弦相似度)
  14. 5. 向量数据入库(测试chroma,主流milvus)
  15. 二. 用户基于知识库提问
  16. 1. 用户提问
  17. 2. 用户问题向量化,到库中检索,top_k返回前n条完整结果
  18. 3. 将结果封装进提示词
  19. 4. 请求LLM返回最终结果
  20. """
  21. from dotenv import load_dotenv
  22. from langchain_openai import ChatOpenAI
  23. import os, re
  24. from fastapi import FastAPI, UploadFile, File, HTTPException
  25. from langchain_community.document_loaders import PyPDFLoader, DirectoryLoader
  26. import tempfile
  27. from langchain_text_splitters import RecursiveCharacterTextSplitter
  28. from langchain_community.vectorstores import Chroma
  29. from langchain_community.embeddings import ZhipuAIEmbeddings
  30. from langchain_core.prompts import PromptTemplate
  31. from zai import ZhipuAiClient
  32. load_dotenv(override=True)
  33. deepseek_base_url = os.getenv("DEEPSEEK_BASE_URL")
  34. deepseek_base_key = os.getenv("DEEPSEEK_BASE_KEY")
  35. deepseek_base_name = os.getenv("DEEPSEEK_BASE_NAME")
  36. # ==============================一. 构建知识库==================================================
  37. # ==============================1. 收集资料==================================================
  38. class KnowledgeBaseBuilder:
  39. def __init__(self, pdf_dir="./resources", persist_dir="./chroma_db"):
  40. self.pdf_dir = pdf_dir
  41. self.persist_dir = persist_dir
  42. self.embeddings = ZhipuAIEmbeddings(
  43. model="embedding-3", # 或者 "embedding-3"
  44. api_key=os.getenv("ZHIPU_API_KEY")
  45. )
  46. self.documents = []
  47. def load_pdfs(self):
  48. """加载PDF文件"""
  49. loader = DirectoryLoader(
  50. self.pdf_dir,
  51. glob="**/*.pdf",
  52. loader_cls=PyPDFLoader,
  53. show_progress=True
  54. )
  55. self.documents = loader.load()
  56. print(f"加载了 {len(self.documents)} 个文档")
  57. return self.documents
  58. # ==============================2. 解析并清洗数据==================================================
  59. def clean_documents(self, docs):
  60. """清洗文档数据"""
  61. cleaned_docs = []
  62. empty_count = 0
  63. for i, doc in enumerate(docs):
  64. text = doc.page_content
  65. # 检查文本是否为空
  66. if not text or len(text.strip()) == 0:
  67. empty_count += 1
  68. print(f"⚠️ 文档 {i+1} 内容为空,跳过")
  69. continue
  70. # ① 清洗噪音:移除页眉页脚、水印等
  71. text = re.sub(r'第\s*\d+\s*页\s*/\s*共\s*\d+\s*页', '', text)
  72. text = re.sub(r'Copyright.*?\n', '', text, flags=re.IGNORECASE)
  73. text = re.sub(r'机密|内部资料|仅供内部使用', '', text)
  74. # ② 去除无意义字符
  75. text = re.sub(r'\n\s*\n+', '\n\n', text) # 合并多个空行
  76. text = re.sub(r'[ \t]+', ' ', text) # 合并多个空格
  77. # 保留中英文、数字和常用标点
  78. text = re.sub(r'[^\u4e00-\u9fa5a-zA-Z0-9\.\,\,\。\!\?\:\;\(\)\n]', ' ', text)
  79. text = re.sub(r'\s+', ' ', text).strip()
  80. # 清洗后再次检查
  81. if not text or len(text) < 10:
  82. empty_count += 1
  83. print(f"⚠️ 文档 {i+1} 清洗后内容过短,跳过")
  84. continue
  85. doc.page_content = text
  86. cleaned_docs.append(doc)
  87. print(f"✅ 清洗完成,有效文档: {len(cleaned_docs)}, 跳过: {empty_count}")
  88. if not cleaned_docs:
  89. raise ValueError("清洗后没有有效的文档内容")
  90. return cleaned_docs
  91. # ==============================3. 分块==================================================
  92. def chunk_documents(self, docs, strategy="recursive"):
  93. """
  94. 多种分块策略
  95. strategy: fixed, recursive, semantic, markdown
  96. 当前实现递归分块策略
  97. """
  98. if strategy == "recursive":
  99. splitter = RecursiveCharacterTextSplitter(
  100. chunk_size=500,
  101. chunk_overlap=50,
  102. separators=[
  103. "\n\n", # 段落
  104. "\n", # 行
  105. "。", # 中文句号
  106. "!", # 感叹号
  107. "?", # 问号
  108. ";", # 分号
  109. ",", # 逗号
  110. ".", # 英文句号
  111. "!",
  112. "?",
  113. " "
  114. ],
  115. length_function=len,
  116. keep_separator=True
  117. )
  118. chunked_docs = splitter.split_documents(docs)
  119. return chunked_docs
  120. # ==============================4. 向量化==================================================
  121. # ==============================5. 向量数据入库==================================================
  122. def build_vectorstore(self, chunks, collection_name="knowledge_base"):
  123. """向量化并存入Chroma数据库"""
  124. # 使用智谱Embedding模型
  125. vectorstore = Chroma.from_documents(
  126. documents=chunks,
  127. embedding=self.embeddings,
  128. persist_directory=self.persist_dir,
  129. collection_name=collection_name
  130. )
  131. vectorstore.persist()
  132. print(f"向量数据库构建完成,存储于: {self.persist_dir}")
  133. return vectorstore
  134. # ==============================二. 用户基于知识库提问==================================================
  135. # ==============================1. 用户提问==================================================
  136. # ==============================2. 用户问题向量化,到库中检索,top_k返回前n条完整结果==================================================
  137. # ==============================3. 将结果封装进提示词==================================================
  138. # ==============================4. 请求LLM返回最终结果==================================================
  139. llm = ChatOpenAI(
  140. base_url = deepseek_base_url,
  141. api_key = deepseek_base_key,
  142. model= deepseek_base_name
  143. )
  144. if __name__ == "__main__":
  145. knowlege = KnowledgeBaseBuilder()
  146. docs = knowlege.load_pdfs()
  147. clean_docs = knowlege.clean_documents(docs)
  148. chunks = knowlege.chunk_documents(clean_docs)
  149. vectorstore = knowlege.build_vectorstore(chunks)
  150. # vectorstore = Chroma(
  151. # persist_directory="./chroma_db", # 持久化目录
  152. # embedding_function=knowlege.embeddings,
  153. # collection_name="knowledge_base" # 集合名称(与构建时一致)
  154. # )
  155. user_question = '遥控器如何操作?'
  156. retrieved_docs = vectorstore.similarity_search(user_question, k=3)
  157. context = "\n\n".join([doc.page_content for doc in retrieved_docs])
  158. prompt_template = """
  159. 你是一个专业的AI助手,请基于以下参考资料回答用户的问题。
  160. 【参考资料】:
  161. {context}
  162. 【用户问题】:
  163. {question}
  164. 要求:
  165. 1. 严格基于参考资料回答,不要编造信息
  166. 2. 如果参考资料中没有相关信息,请明确告知
  167. 3. 回答要简洁、准确、有条理
  168. 回答:
  169. """
  170. prompt = PromptTemplate(
  171. template=prompt_template,
  172. input_variables=["context", "question"]
  173. )
  174. # 4. 请求LLM
  175. formatted_prompt = prompt.format(context=context, question=user_question)
  176. response = llm.invoke(formatted_prompt)
  177. print(response)