02_RAG_task.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543
  1. """
  2. RAG系统实战 - 投资方法论知识库
  3. 功能:加载PDF文档、数据清洗、分块、向量化、存储到向量库、检索生成回答
  4. """
  5. import os
  6. import re
  7. from typing import List
  8. from dotenv import load_dotenv
  9. # LangChain 核心组件
  10. from langchain_community.document_loaders import PyMuPDFLoader
  11. from langchain_text_splitters import RecursiveCharacterTextSplitter
  12. from langchain_community.embeddings import DashScopeEmbeddings
  13. from langchain_community.vectorstores import Chroma
  14. from langchain_core.prompts import ChatPromptTemplate
  15. from langchain_openai import ChatOpenAI
  16. from langchain_core.output_parsers import StrOutputParser
  17. from langchain_core.documents import Document
  18. # 加载环境变量
  19. load_dotenv()
  20. class RAGSystem:
  21. """RAG系统类,封装完整的检索增强生成流程"""
  22. def __init__(self, pdf_path: str, persist_directory: str = "./knowledge_db"):
  23. """
  24. 初始化RAG系统
  25. Args:
  26. pdf_path: PDF文件路径
  27. persist_directory: 向量数据库持久化目录
  28. """
  29. self.pdf_path = pdf_path
  30. self.persist_directory = persist_directory
  31. self.pages = None
  32. self.split_docs = None
  33. self.vectorstore = None
  34. self.retriever = None
  35. # 初始化Embedding模型(使用阿里通义)
  36. self.embedding_model = DashScopeEmbeddings(
  37. model=os.getenv("EMBEDDING_MODEL", "text-embedding-v3"),
  38. dashscope_api_key=os.getenv("DASHSCOPE_API_KEY", "")
  39. )
  40. # 初始化大模型(使用DeepSeek)
  41. self.llm = ChatOpenAI(
  42. model=os.getenv("MODEL_NAME", "deepseek-chat"),
  43. openai_api_key=os.getenv("OPENAI_API_KEY", ""),
  44. openai_api_base=os.getenv("OPENAI_API_BASE", "https://api.deepseek.com/v1"),
  45. temperature=0.7
  46. )
  47. log_info("RAG系统初始化完成")
  48. log_info(f"使用模型: {os.getenv('MODEL_NAME', 'deepseek-chat')}")
  49. log_info(f"API地址: {os.getenv('OPENAI_API_BASE', 'https://api.deepseek.com/v1')}")
  50. def load_pdf(self) -> List[Document]:
  51. """
  52. 步骤1:加载PDF文档
  53. Returns:
  54. 文档页面列表
  55. """
  56. log_info(f"正在加载PDF文档: {self.pdf_path}")
  57. # 检查文件是否存在
  58. if not os.path.exists(self.pdf_path):
  59. raise FileNotFoundError(f"PDF文件不存在: {self.pdf_path}")
  60. # 使用PyMuPDF加载PDF
  61. loader = PyMuPDFLoader(self.pdf_path)
  62. self.pages = loader.load()
  63. log_info(f"PDF加载完成,共 {len(self.pages)} 页")
  64. # 检查是否有实际内容
  65. total_content = sum(len(page.page_content) for page in self.pages)
  66. if total_content == 0:
  67. log_error("PDF文件没有可提取的文本内容!")
  68. log_error("可能原因:")
  69. log_error(" 1. PDF是扫描版或图片型PDF")
  70. log_error(" 2. PDF有加密保护")
  71. log_error(" 3. PDF格式特殊")
  72. log_error("\n解决方案:")
  73. log_error(" 1. 使用OCR工具将PDF转换为文本")
  74. log_error(" 2. 尝试其他PDF文件")
  75. raise ValueError("PDF文件无可提取文本内容,请使用文本型PDF或进行OCR处理")
  76. # 打印统计信息
  77. non_empty_pages = sum(1 for page in self.pages if len(page.page_content) > 0)
  78. log_info(f"有内容的页数: {non_empty_pages}/{len(self.pages)}")
  79. log_info(f"总字符数: {total_content}")
  80. # 打印第一页预览
  81. if self.pages and len(self.pages[0].page_content) > 0:
  82. preview = self.pages[0].page_content[:500]
  83. log_info(f"第一页内容预览:\n{preview}...")
  84. return self.pages
  85. def clean_text(self, text: str) -> str:
  86. """
  87. 步骤2:清洗PDF解析出的文本,去除常见噪声
  88. Args:
  89. text: 原始文本
  90. Returns:
  91. 清洗后的文本
  92. """
  93. if not text or not text.strip():
  94. return ""
  95. # 保存原始长度用于对比
  96. original_len = len(text)
  97. # 删除连续的换行符(保留最多2个)
  98. text = re.sub(r'\n{3,}', '\n\n', text)
  99. # 删除项目符号
  100. text = text.replace('•', '').replace('·', '')
  101. # 合并多余空格(但不删除换行符周围的空格)
  102. text = re.sub(r'[^\S\n]+', ' ', text)
  103. # 删除特殊控制字符
  104. text = re.sub(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f-\x9f]', '', text)
  105. # 如果清洗后文本为空或太短,返回原始文本
  106. cleaned_text = text.strip()
  107. if len(cleaned_text) < 10: # 如果清洗后少于10个字符,返回原始文本
  108. log_info(f"警告: 清洗后文本过短({len(cleaned_text)}字符),保留原始文本({original_len}字符)")
  109. return text.strip()
  110. return cleaned_text
  111. def clean_all_pages(self) -> List[Document]:
  112. """
  113. 清洗所有页面
  114. Returns:
  115. 清洗后的文档列表
  116. """
  117. if not self.pages:
  118. raise ValueError("请先调用 load_pdf() 加载文档")
  119. log_info("开始清洗文档...")
  120. # 清洗每一页的内容
  121. for page in self.pages:
  122. page.page_content = self.clean_text(page.page_content)
  123. log_info(f"文档清洗完成,共处理 {len(self.pages)} 页")
  124. return self.pages
  125. def split_documents(
  126. self,
  127. chunk_size: int = 100,
  128. chunk_overlap: int = 20
  129. ) -> List[Document]:
  130. """
  131. 步骤3:文档分块(使用递归字符分割器)
  132. Args:
  133. chunk_size: 每个块的最大字符数
  134. chunk_overlap: 相邻块之间的重叠字符数
  135. Returns:
  136. 分块后的文档列表
  137. """
  138. if not self.pages:
  139. raise ValueError("请先调用 load_pdf() 加载文档")
  140. log_info(f"开始分块,chunk_size={chunk_size}, chunk_overlap={chunk_overlap}")
  141. # 创建递归字符分割器
  142. text_splitter = RecursiveCharacterTextSplitter(
  143. # 分隔符优先级:段落 → 换行 → 句号 → 空格 → 硬切
  144. separators=["\n\n", "\n", "。", "!", "?", ";", " ", ""],
  145. chunk_size=chunk_size,
  146. chunk_overlap=chunk_overlap,
  147. length_function=len,
  148. is_separator_regex=False
  149. )
  150. # 执行分块
  151. self.split_docs = text_splitter.split_documents(self.pages)
  152. # 过滤空块
  153. self.split_docs = [
  154. doc for doc in self.split_docs
  155. if doc.page_content and doc.page_content.strip()
  156. ]
  157. # 统计信息
  158. total_chars = sum(len(doc.page_content) for doc in self.split_docs)
  159. log_info(f"分块完成,共 {len(self.split_docs)} 个块,总字符数: {total_chars}")
  160. # 打印前3个块预览
  161. for i, doc in enumerate(self.split_docs[:3]):
  162. log_info(f"块{i+1} 预览: {doc.page_content[:100]}...")
  163. return self.split_docs
  164. def create_vectorstore(self) -> Chroma:
  165. """
  166. 步骤4:创建向量数据库并存储文档
  167. Returns:
  168. 向量数据库实例
  169. """
  170. if not self.split_docs:
  171. raise ValueError("请先调用 split_documents() 分块文档")
  172. log_info("开始创建向量数据库...")
  173. # 创建向量数据库并持久化
  174. self.vectorstore = Chroma.from_documents(
  175. documents=self.split_docs,
  176. embedding=self.embedding_model,
  177. collection_metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
  178. persist_directory=self.persist_directory
  179. )
  180. log_info(f"向量数据库创建完成,保存至: {self.persist_directory}")
  181. return self.vectorstore
  182. def load_vectorstore(self) -> Chroma:
  183. """
  184. 加载已存在的向量数据库
  185. Returns:
  186. 向量数据库实例
  187. """
  188. if not os.path.exists(self.persist_directory):
  189. raise FileNotFoundError(f"向量数据库不存在: {self.persist_directory}")
  190. log_info(f"加载向量数据库: {self.persist_directory}")
  191. self.vectorstore = Chroma(
  192. persist_directory=self.persist_directory,
  193. embedding_function=self.embedding_model
  194. )
  195. log_info("向量数据库加载完成")
  196. return self.vectorstore
  197. def create_retriever(self, k: int = 3):
  198. """
  199. 步骤5:创建检索器
  200. Args:
  201. k: 返回的最相关文档数量
  202. Returns:
  203. 检索器实例
  204. """
  205. if not self.vectorstore:
  206. raise ValueError("请先创建或加载向量数据库")
  207. log_info(f"创建检索器,返回Top-{k}文档")
  208. self.retriever = self.vectorstore.as_retriever(
  209. search_kwargs={"k": k}
  210. )
  211. return self.retriever
  212. def retrieve(self, query: str, k: int = 3) -> List[Document]:
  213. """
  214. 步骤6:检索相关文档
  215. Args:
  216. query: 用户问题
  217. k: 返回的文档数量
  218. Returns:
  219. 相关文档列表
  220. """
  221. if not self.retriever:
  222. raise ValueError("请先创建检索器")
  223. log_info(f"检索问题: {query}")
  224. # 执行检索
  225. relevant_docs = self.retriever.invoke(query)
  226. log_info(f"检索完成,找到 {len(relevant_docs)} 个相关文档")
  227. # 打印检索结果
  228. for i, doc in enumerate(relevant_docs):
  229. log_info(f"文档{i+1}: {doc.page_content[:150]}...")
  230. return relevant_docs
  231. def generate_answer(self, query: str, relevant_docs: List[Document]) -> str:
  232. """
  233. 步骤7:基于检索结果生成回答
  234. Args:
  235. query: 用户问题
  236. relevant_docs: 相关文档列表
  237. Returns:
  238. 生成的回答
  239. """
  240. log_info("开始生成回答...")
  241. # 拼接上下文
  242. context = "\n\n---\n\n".join([doc.page_content for doc in relevant_docs])
  243. # 构建Prompt模板
  244. prompt = ChatPromptTemplate.from_template("""
  245. 你是一个专业的投资知识库助手。请根据以下检索到的上下文回答用户问题。
  246. **规则:**
  247. - 只基于提供的上下文回答,不要编造
  248. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  249. - 回答要简洁直接,引用原文时用引号
  250. - 回答时请标注信息来源的页码
  251. **检索到的上下文:**
  252. {context}
  253. **用户问题:**
  254. {question}
  255. """)
  256. # 构建Chain
  257. chain = prompt | self.llm | StrOutputParser()
  258. # 生成回答
  259. answer = chain.invoke({
  260. "context": context,
  261. "question": query
  262. })
  263. log_info("回答生成完成")
  264. return answer
  265. def query(self, question: str, k: int = 3) -> str:
  266. """
  267. 完整查询流程:检索 + 生成
  268. Args:
  269. question: 用户问题
  270. k: 检索文档数量
  271. Returns:
  272. 生成的回答
  273. """
  274. # 检索相关文档
  275. relevant_docs = self.retrieve(question, k)
  276. # 生成回答
  277. answer = self.generate_answer(question, relevant_docs)
  278. return answer
  279. def build_knowledge_base(self):
  280. """
  281. 构建完整的知识库:加载 → 清洗 → 分块 → 向量化 → 存储
  282. """
  283. log_info("=" * 50)
  284. log_info("开始构建知识库...")
  285. log_info("=" * 50)
  286. # 步骤1:加载PDF
  287. self.load_pdf()
  288. # 步骤2:清洗数据
  289. self.clean_all_pages()
  290. # 步骤3:分块
  291. self.split_documents()
  292. # 步骤4:向量化并存储
  293. self.create_vectorstore()
  294. # 步骤5:创建检索器
  295. self.create_retriever()
  296. log_info("=" * 50)
  297. log_info("知识库构建完成!")
  298. log_info("=" * 50)
  299. def log_info(message: str):
  300. """打印日志信息"""
  301. print(f"[INFO] {message}")
  302. def interactive_query(rag_system: RAGSystem):
  303. """
  304. 交互式问答模式
  305. Args:
  306. rag_system: RAG系统实例
  307. """
  308. print("\n" + "=" * 60)
  309. print("RAG知识库问答系统 - 投资方法论")
  310. print("=" * 60)
  311. print("输入问题开始查询,输入 'quit' 或 'exit' 退出\n")
  312. while True:
  313. try:
  314. # 获取用户输入
  315. question = input("你的问题: ").strip()
  316. # 检查退出命令
  317. if question.lower() in ['quit', 'exit', 'q']:
  318. print("\n感谢使用,再见!")
  319. break
  320. # 跳过空问题
  321. if not question:
  322. print("请输入有效问题\n")
  323. continue
  324. # 执行查询
  325. print("\n正在检索并生成答案...\n")
  326. answer = rag_system.query(question)
  327. # 显示结果
  328. print("-" * 60)
  329. print(f"回答:\n{answer}")
  330. print("-" * 60)
  331. print()
  332. except KeyboardInterrupt:
  333. print("\n\n感谢使用,再见!")
  334. break
  335. except Exception as e:
  336. log_error(f"查询出错: {str(e)}")
  337. def log_error(message: str):
  338. """打印错误日志"""
  339. print(f"[ERROR] {message}")
  340. def main():
  341. """主函数"""
  342. # PDF文件路径
  343. pdf_path = r"D:\investment\data\疯狂的里海_投资体系框架.pdf"
  344. # 备用PDF文件路径(用于测试)
  345. backup_pdf_path = r"./car_info.pdf"
  346. # 向量数据库存储路径
  347. persist_directory = "D:\agentlearning\lqq-agent-study\investment_db"
  348. # 检查API Key配置
  349. openai_api_key = os.getenv("OPENAI_API_KEY")
  350. dashscope_api_key = os.getenv("DASHSCOPE_API_KEY")
  351. if not openai_api_key:
  352. print("警告: 未设置 OPENAI_API_KEY 环境变量")
  353. print("请在 .env 文件中添加: OPENAI_API_KEY=your-deepseek-api-key")
  354. print("或者访问 https://platform.deepseek.com/ 获取API Key")
  355. return
  356. if not dashscope_api_key:
  357. print("警告: 未设置 DASHSCOPE_API_KEY 环境变量")
  358. print("请在 .env 文件中添加: DASHSCOPE_API_KEY=your-dashscope-api-key")
  359. print("或者访问 https://dashscope.console.aliyun.com/ 获取API Key")
  360. return
  361. try:
  362. # 检查主PDF文件是否存在
  363. if not os.path.exists(pdf_path):
  364. print(f"警告: 指定的PDF文件不存在: {pdf_path}")
  365. # 尝试使用备用PDF文件
  366. if os.path.exists(backup_pdf_path):
  367. print(f"将使用备用PDF文件进行测试: {backup_pdf_path}")
  368. pdf_path = backup_pdf_path
  369. persist_directory = "./car_info_knowledge_db"
  370. else:
  371. print("备用PDF文件也不存在,请检查文件路径")
  372. return
  373. # 创建RAG系统实例
  374. rag = RAGSystem(
  375. pdf_path=pdf_path,
  376. persist_directory=persist_directory
  377. )
  378. # 检查向量数据库是否已存在
  379. if os.path.exists(persist_directory):
  380. print(f"检测到已有向量数据库: {persist_directory}")
  381. choice = input("是否重新构建知识库?(y/n): ").strip().lower()
  382. if choice == 'y':
  383. # 重新构建知识库
  384. rag.build_knowledge_base()
  385. else:
  386. # 加载已有向量数据库
  387. rag.load_vectorstore()
  388. rag.create_retriever()
  389. else:
  390. # 构建新知识库
  391. rag.build_knowledge_base()
  392. # 进入交互式问答
  393. interactive_query(rag)
  394. except ValueError as e:
  395. if "PDF文件无可提取文本内容" in str(e):
  396. print("\n" + "=" * 60)
  397. print("PDF文件处理失败")
  398. print("=" * 60)
  399. print("\n您的PDF文件是扫描版或图片型PDF,无法直接提取文本。")
  400. print("\n解决方案:")
  401. print("1. 使用OCR工具将PDF转换为文本")
  402. print(" 推荐工具:Adobe Acrobat、福昕PDF、ABBYY FineReader")
  403. print(" 在线工具:https://www.pdf2go.com/zh/ocr-pdf")
  404. print("\n2. 使用项目中的示例PDF文件进行测试:")
  405. print(f" 文件路径: {backup_pdf_path}")
  406. print("\n3. 寻找其他文本型PDF文件")
  407. print("\n提示:您可以修改代码中的pdf_path变量指向可用的PDF文件")
  408. else:
  409. log_error(str(e))
  410. except FileNotFoundError as e:
  411. log_error(str(e))
  412. log_error("请检查PDF文件路径是否正确")
  413. except Exception as e:
  414. log_error(f"系统运行出错: {str(e)}")
  415. import traceback
  416. traceback.print_exc()
  417. if __name__ == "__main__":
  418. main()