rag_demo.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146
  1. from langchain_community.document_loaders import PyMuPDFLoader
  2. from langchain_text_splitters import RecursiveCharacterTextSplitter
  3. from langchain_community.embeddings import DashScopeEmbeddings
  4. from langchain_community.vectorstores import Chroma
  5. from langchain_core.prompts import ChatPromptTemplate
  6. from langchain_community.chat_models import ChatTongyi
  7. from langchain_core.output_parsers import StrOutputParser
  8. from dotenv import load_dotenv
  9. import re
  10. import os
  11. load_dotenv()
  12. api_key = os.getenv('QWEN_API_KEY')
  13. # ============================================================
  14. # 1. 加载 PDF 文档
  15. # ============================================================
  16. loader = PyMuPDFLoader("D:/code/shuheAI/02_RAG/华为擎云W585X 用户指南-(PGUX,KOS&UOS_02,zh-cn).pdf")
  17. pdf_pages = loader.load()
  18. print(f"文档类型:{type(pdf_pages)}")
  19. print(f"PDF 共{len(pdf_pages)}页")
  20. # 查看第一页的内容和元数据
  21. first_page = pdf_pages[0]
  22. print(f"元数据:{first_page.metadata}")
  23. print(f"内容预览:{first_page.page_content[:200]}")
  24. # ============================================================
  25. # 2. 清洗 PDF 文本
  26. # ============================================================
  27. def clean_pdf_text(text: str) -> str:
  28. """清洗 PDF 解析出的文本,去除常见噪声"""
  29. # 删除非中文字符之间的换行符
  30. text = re.sub(r'[^一](\n)[^一]', lambda m: m.group(0).replace('\n', ''), text)
  31. # 删除项目符号和多余空格
  32. text = text.replace('•', '').replace(' ', ' ')
  33. # 删除连续的换行符(保留一个)
  34. text = re.sub(r'\n{2,}', '\n', text)
  35. return text.strip()
  36. # 对所有页面清洗
  37. for page in pdf_pages:
  38. page.page_content = clean_pdf_text(page.page_content)
  39. # ============================================================
  40. # 3. 文本分块
  41. # ============================================================
  42. splitter = RecursiveCharacterTextSplitter(
  43. separators=["\n\n", "\n", "。", "!", "?", " ", ""],
  44. chunk_size=50,
  45. chunk_overlap=10,
  46. length_function=len
  47. )
  48. split_docs = splitter.split_documents(pdf_pages)
  49. print(f"切分后的文件数量:{len(split_docs)}")
  50. print(f"切分后的字符数(可以用来大致评估 token 数):{sum([len(doc.page_content) for doc in split_docs])}")
  51. # 过滤掉 page_content 为空或仅含空白的文档
  52. valid_docs = [
  53. doc for doc in split_docs
  54. if doc.page_content and doc.page_content.strip()
  55. ]
  56. print(f"有效块数量:{len(valid_docs)}")
  57. print(f"总字符数(可大致评估 Token 数):{sum(len(d.page_content) for d in valid_docs)}")
  58. # ============================================================
  59. # 4. 初始化 Embedding 模型
  60. # ============================================================
  61. embedding_model = DashScopeEmbeddings(
  62. model="text-embedding-v3",
  63. dashscope_api_key=api_key
  64. )
  65. # 单条文本向量化
  66. text = "RAG系统搭建实战"
  67. embedding = embedding_model.embed_query(text)
  68. print(f"向量维度:{len(embedding)}")
  69. print(f"前5个值:{embedding[:5]}")
  70. # ============================================================
  71. # 5. 向量化 + 存入向量库
  72. # ============================================================
  73. persist_dir = "./my_knowledge_db"
  74. if os.path.exists(persist_dir) and os.listdir(persist_dir):
  75. print(f"加载已有向量库:{persist_dir}")
  76. vectordb = Chroma(
  77. persist_directory=persist_dir,
  78. embedding_function=embedding_model
  79. )
  80. else:
  81. print(f"创建新向量库:{persist_dir}")
  82. vectordb = Chroma.from_documents(
  83. documents=valid_docs,
  84. embedding=embedding_model,
  85. collection_metadata={"hnsw:space": "cosine"}, # 余弦相似度
  86. persist_directory=persist_dir
  87. )
  88. # 创建检索器,设置返回 Top-3 最相关文档
  89. retriever = vectordb.as_retriever(search_kwargs={"k": 3})
  90. # ============================================================
  91. # 6. 提问
  92. # ============================================================
  93. query = "如何进入BIOS设置?"
  94. relevant_docs = retriever.invoke(query)
  95. for i, doc in enumerate(relevant_docs):
  96. print(f"--- 结果{i+1} ---")
  97. print(f"内容:{doc.page_content[:100]}...")
  98. print(f"来源:{doc.metadata}")
  99. print()
  100. # 带分数的相似度检索(分数越低越相似,0 表示完全匹配)
  101. results = vectordb.similarity_search_with_score(query, k=3)
  102. for doc, score in results:
  103. print(f"内容:{doc.page_content[:100]}... | 相似度分数:{score:.4f}")
  104. # ============================================================
  105. # 7. 生成回答
  106. # ============================================================
  107. context = "\n\n---\n\n".join([d.page_content for d in relevant_docs])
  108. prompt = ChatPromptTemplate.from_template("""
  109. 你是一个专业的知识库助手。请根据以下上下文回答问题。
  110. **规则:**
  111. - 只基于提供的上下文回答,不要编造
  112. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  113. - 回答要简洁直接,引用原文时用引号
  114. **上下文:**
  115. {context}
  116. **问题:**
  117. {question}
  118. """)
  119. llm = ChatTongyi(model="qwen-plus", dashscope_api_key=api_key)
  120. chain = prompt | llm | StrOutputParser()
  121. answer = chain.invoke({"context": context, "question": query})
  122. print(answer)