rag_system.py 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131
  1. import os
  2. from dotenv import load_dotenv
  3. from langchain_community.document_loaders import PyMuPDFLoader
  4. from langchain_text_splitters import RecursiveCharacterTextSplitter
  5. from langchain_openai import OpenAIEmbeddings
  6. # from langchain_community.vectorstores import Chroma
  7. from langchain_chroma import Chroma
  8. from langchain_core.prompts import ChatPromptTemplate
  9. from langchain_community.chat_models import ChatTongyi
  10. from langchain_core.output_parsers import StrOutputParser
  11. from langchain_openai import ChatOpenAI
  12. load_dotenv(dotenv_path=".env")
  13. emb_base_url = os.getenv("EMB_MODEL_URL")
  14. emb_model_name = os.getenv("EMB_MODEL_NAME")
  15. chat_base_url = os.getenv("CHAT_MODEL_URL")
  16. chat_model_name = os.getenv("CHAT_MODEL_NAME")
  17. # ========== 第一步:加载文档 ==========
  18. loader = PyMuPDFLoader("中国近代史.pdf")
  19. pages = loader.load()
  20. # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
  21. # clean_pages = [clean_pdf_text(page) for page in pages]
  22. # ========== 第三步:分块 ==========
  23. # splitter = RecursiveCharacterTextSplitter(
  24. # # 分隔符优先级:段落 → 换行 → 句号 → 空格 → 硬切
  25. # separators=["\n\n", "\n", "。", "!", "?", " ", ""],
  26. # # 每个块最大 50 字符
  27. # chunk_size=180,
  28. # # 相邻块重叠 10 字符(chunk_size 的 20%)
  29. # chunk_overlap=10,
  30. # # 长度计算函数
  31. # length_function=len
  32. # )
  33. # docs = splitter.split_documents(pages)
  34. # ========== 第四步:向量化 + 存入向量库 ==========
  35. embedding_model = OpenAIEmbeddings(
  36. model=emb_model_name, # 本地模型名称(根据实际部署填写)
  37. api_key="not-needed", # 本地服务通常不需要真实 key,填任意值即可
  38. base_url=emb_base_url # 本地服务地址,根据实际端口修改
  39. )
  40. # embedding_model = OpenAIEmbeddings(
  41. # model="BAAI/bge-m3", # 本地模型名称(根据实际部署填写)
  42. # api_key="sk-hynxocghcwhevjapdsccungulfggjlaphlxiyenghcotcwux", # 本地服务通常不需要真实 key,填任意值即可
  43. # base_url="https://api.siliconflow.cn/v1" # 本地服务地址,根据实际端口修改
  44. # )
  45. #
  46. # vectorstore = Chroma.from_documents(
  47. # documents=docs,
  48. # embedding=embedding_model,
  49. # collection_metadata={"hnsw:space": "cosine"},
  50. # persist_directory="./knowledge_db"
  51. # )
  52. #
  53. # exit()
  54. # ========== 加载已有向量库 ==========
  55. vectorstore = Chroma(
  56. persist_directory="./knowledge_db", # 之前保存的目录
  57. embedding_function=embedding_model # 加载时也需要指定 embedding 模型
  58. )
  59. # ========== 查看向量库前十条chunk ==========
  60. all_data = vectorstore.get(limit=10)
  61. # 查看结构
  62. # print(all_data.keys())
  63. # 输出: dict_keys(['ids', 'embeddings', 'documents', 'metadatas'])
  64. # 查看所有 chunk 的文本内容
  65. for i, doc in enumerate(all_data['documents']):
  66. print(f"\n--- Chunk {i} ---")
  67. print(f"ID: {all_data['ids'][i]}")
  68. print(f"文本: {doc}")
  69. # print(f"元数据: {all_data['metadatas'][i]}")
  70. # # ========== 第五步:创建检索器 ==========
  71. # retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
  72. # # ========== 第六步:提问 ==========
  73. query = "徐中约是谁"
  74. relevant_docs = vectorstore.similarity_search_with_score(query, k=3)
  75. for doc, score in relevant_docs:
  76. print(f"{doc.page_content}")
  77. print("-"* 30)
  78. # ========== 第七步:生成回答 ==========
  79. context = "\n\n---\n\n".join([content.page_content for content, score in relevant_docs])
  80. # print(context)
  81. #
  82. prompt = ChatPromptTemplate.from_template("""
  83. 你是一个专业的知识库助手。请根据以下上下文回答问题。
  84. **规则:**
  85. - 只基于提供的上下文回答,不要编造
  86. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  87. - 回答要简洁直接,引用原文时用引号
  88. **上下文:**
  89. {context}
  90. **问题:**
  91. {question}
  92. """)
  93. #
  94. llm = ChatOpenAI(
  95. model=chat_model_name,
  96. api_key="api_key",
  97. base_url=chat_base_url,
  98. temperature=0.5,
  99. streaming=True
  100. )
  101. chain = prompt | llm | StrOutputParser()
  102. # print(chain)
  103. #
  104. # answer = chain.invoke({"context": context, "question": query})
  105. for chunk in chain.stream({"context": context, "question": query}):
  106. print(chunk, end="", flush=True)
  107. # print(answer)