main.py 2.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990
  1. import os
  2. from langchain_community.document_loaders import PyMuPDFLoader
  3. from langchain_text_splitters import RecursiveCharacterTextSplitter
  4. from utils.dataclean import clean_pdf_text
  5. from langchain_community.vectorstores import Chroma
  6. from langchain_core.prompts import ChatPromptTemplate
  7. from langchain_core.output_parsers import StrOutputParser
  8. from config import embedding_model, ChatTongyillm
  9. # ========== 第一步:加载文档 ==========
  10. # 获取当前脚本所在目录,构建 PDF 文件的绝对路径
  11. script_dir = os.path.dirname(os.path.abspath(__file__))
  12. pdf_path = os.path.join(script_dir, "data", "car_info.pdf")
  13. # 创建加载器实例,传入 PDF 文件路径
  14. pdf_loader = PyMuPDFLoader(pdf_path)
  15. # 调用 load() 方法,返回一个 Document 列表(每页一个 Document)
  16. pdf_pages = pdf_loader.load()
  17. print("加载资料完成")
  18. # ========== 第二步:清洗数据(可选,根据文档质量决定)==========
  19. # 清洗每个 Document 的文本内容
  20. for page in pdf_pages:
  21. page.page_content = clean_pdf_text(page.page_content)
  22. print("清晰资料完成")
  23. # ========== 第三步:创建递归字符分割器 ==========
  24. text_splitter = RecursiveCharacterTextSplitter(
  25. # 分隔符优先级:段落 → 换行 → 句号 → 空格 → 硬切
  26. separators=["\n\n", "\n", "。", "!", "?", " ", ""],
  27. # 每个块最大 50 字符
  28. chunk_size=50,
  29. # 相邻块重叠 10 字符(chunk_size 的 20%)
  30. chunk_overlap=10,
  31. # 长度计算函数
  32. length_function=len
  33. )
  34. # 对整个文档切分
  35. split_docs = text_splitter.split_documents(pdf_pages)
  36. print("切割资料完成")
  37. # ========== 第四步:向量化 + 存入向量库 ==========
  38. vectorstore = Chroma.from_documents(
  39. documents=split_docs,
  40. embedding=embedding_model,
  41. persist_directory="./knowledge_db"
  42. )
  43. print("存储向量资料完成")
  44. # ========== 第五步:创建检索器 ==========
  45. retriever = vectorstore.as_retriever(search_kwargs={"k": 3})
  46. # ========== 第六步:提问 ==========
  47. query = "什么汽车比较经济实惠"
  48. relevant_docs = retriever.invoke(query)
  49. # ========== 第七步:生成回答 ==========
  50. context = "\n\n---\n\n".join([d.page_content for d in relevant_docs])
  51. prompt = ChatPromptTemplate.from_template("""
  52. 你是一个专业的知识库助手。请根据以下上下文回答问题。
  53. **规则:**
  54. - 只基于提供的上下文回答,不要编造
  55. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  56. - 回答要简洁直接,引用原文时用引号
  57. **上下文:**
  58. {context}
  59. **问题:**
  60. {question}
  61. """)
  62. chain = prompt | ChatTongyillm | StrOutputParser()
  63. if __name__ == "__main__":
  64. answer = chain.invoke({"context": context, "question": query})
  65. print(answer)