quick_start.py 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188
  1. """
  2. 快速开始示例 - RAG系统
  3. 演示如何使用RAG系统进行知识库问答
  4. """
  5. import os
  6. import sys
  7. from dotenv import load_dotenv
  8. # 添加当前目录到系统路径
  9. sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
  10. # 导入RAG系统(使用importlib处理数字开头的文件名)
  11. import importlib.util
  12. spec = importlib.util.spec_from_file_location("rag_task", "./02_RAG_task.py")
  13. rag_module = importlib.util.module_from_spec(spec)
  14. spec.loader.exec_module(rag_module)
  15. RAGSystem = rag_module.RAGSystem
  16. # 加载环境变量
  17. load_dotenv()
  18. def example_1_build_and_query():
  19. """示例1:构建知识库并查询"""
  20. print("\n=== 示例1:构建知识库并查询 ===\n")
  21. # 创建RAG系统
  22. rag = RAGSystem(
  23. pdf_path=r"D:\investment\疯狂的里海 · 投资方法论 — 基于 367 篇投资周记提炼.pdf",
  24. persist_directory="./investment_knowledge_db"
  25. )
  26. # 构建知识库(首次运行需要,后续可跳过)
  27. rag.build_knowledge_base()
  28. # 提问
  29. questions = [
  30. "什么是价值投资?",
  31. "投资的核心原则是什么?",
  32. "如何判断一个公司是否值得投资?"
  33. ]
  34. for question in questions:
  35. print(f"\n问题: {question}")
  36. answer = rag.query(question)
  37. print(f"回答: {answer}")
  38. print("-" * 60)
  39. def example_2_load_and_query():
  40. """示例2:加载已有知识库并查询"""
  41. print("\n=== 示例2:加载已有知识库并查询 ===\n")
  42. # 创建RAG系统
  43. rag = RAGSystem(
  44. pdf_path=r"D:\investment\疯狂的里海 · 投资方法论 — 基于 367 篇投资周记提炼.pdf",
  45. persist_directory="./investment_knowledge_db"
  46. )
  47. # 加载已有向量数据库
  48. rag.load_vectorstore()
  49. rag.create_retriever()
  50. # 单次查询
  51. question = "投资者应该如何控制风险?"
  52. print(f"问题: {question}")
  53. answer = rag.query(question)
  54. print(f"回答: {answer}")
  55. def example_3_custom_retrieval():
  56. """示例3:自定义检索参数"""
  57. print("\n=== 示例3:自定义检索参数 ===\n")
  58. # 创建RAG系统
  59. rag = RAGSystem(
  60. pdf_path=r"D:\investment\疯狂的里海 · 投资方法论 — 基于 367 篇投资周记提炼.pdf",
  61. persist_directory="./investment_knowledge_db"
  62. )
  63. # 加载知识库
  64. rag.load_vectorstore()
  65. # 自定义检索器:返回更多相关文档
  66. rag.create_retriever(k=5)
  67. # 检索相关文档
  68. question = "投资周记中提到了哪些投资策略?"
  69. relevant_docs = rag.retrieve(question, k=5)
  70. print(f"问题: {question}")
  71. print(f"\n找到 {len(relevant_docs)} 个相关文档片段:\n")
  72. for i, doc in enumerate(relevant_docs, 1):
  73. print(f"--- 文档片段 {i} ---")
  74. print(f"内容: {doc.page_content[:200]}...")
  75. print(f"来源: {doc.metadata}")
  76. print()
  77. def example_4_step_by_step():
  78. """示例4:分步骤执行RAG流程"""
  79. print("\n=== 示例4:分步骤执行RAG流程 ===\n")
  80. # 创建RAG系统
  81. rag = RAGSystem(
  82. pdf_path=r"D:\investment\疯狂的里海 · 投资方法论 — 基于 367 篇投资周记提炼.pdf",
  83. persist_directory="./custom_knowledge_db"
  84. )
  85. # 步骤1:加载PDF
  86. print("步骤1: 加载PDF文档...")
  87. pages = rag.load_pdf()
  88. print(f" 加载了 {len(pages)} 页\n")
  89. # 步骤2:清洗数据
  90. print("步骤2: 清洗数据...")
  91. rag.clean_all_pages()
  92. print(" 数据清洗完成\n")
  93. # 步骤3:分块(自定义参数)
  94. print("步骤3: 文档分块...")
  95. docs = rag.split_documents(chunk_size=1000, chunk_overlap=200)
  96. print(f" 分成 {len(docs)} 个块\n")
  97. # 步骤4:创建向量库
  98. print("步骤4: 创建向量数据库...")
  99. vectorstore = rag.create_vectorstore()
  100. print(" 向量数据库创建完成\n")
  101. # 步骤5:创建检索器
  102. print("步骤5: 创建检索器...")
  103. retriever = rag.create_retriever(k=3)
  104. print(" 检索器创建完成\n")
  105. # 步骤6:检索
  106. print("步骤6: 检索相关文档...")
  107. question = "什么是安全边际?"
  108. relevant_docs = rag.retrieve(question)
  109. print(f" 找到 {len(relevant_docs)} 个相关文档\n")
  110. # 步骤7:生成答案
  111. print("步骤7: 生成答案...")
  112. answer = rag.generate_answer(question, relevant_docs)
  113. print(f"问题: {question}")
  114. print(f"回答: {answer}")
  115. def main():
  116. """主函数"""
  117. # 检查API Key
  118. if not os.getenv("DASHSCOPE_API_KEY") or os.getenv("DASHSCOPE_API_KEY") == "your_dashscope_api_key_here":
  119. print("错误: 请先在 .env 文件中配置 DASHSCOPE_API_KEY")
  120. print("获取API Key: https://dashscope.console.aliyun.com/")
  121. return
  122. print("=" * 60)
  123. print("RAG系统快速开始示例")
  124. print("=" * 60)
  125. # 选择要运行的示例
  126. print("\n请选择示例:")
  127. print("1. 构建知识库并查询(首次运行)")
  128. print("2. 加载已有知识库并查询")
  129. print("3. 自定义检索参数")
  130. print("4. 分步骤执行RAG流程")
  131. print("0. 退出")
  132. choice = input("\n请输入选项 (0-4): ").strip()
  133. if choice == "1":
  134. example_1_build_and_query()
  135. elif choice == "2":
  136. example_2_load_and_query()
  137. elif choice == "3":
  138. example_3_custom_retrieval()
  139. elif choice == "4":
  140. example_4_step_by_step()
  141. elif choice == "0":
  142. print("退出程序")
  143. else:
  144. print("无效选项")
  145. if __name__ == "__main__":
  146. main()