rag_chain.py 3.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081
  1. from dotenv import load_dotenv # 加载 .env 文件中的环境变量
  2. from pathlib import Path
  3. load_dotenv(Path(__file__).parent.parent / '.env') # 用脚本文件的绝对路径定位 .env
  4. from langchain_community.embeddings import DashScopeEmbeddings # 阿里 DashScope 文本向量化模型,用于把文本 chunks 转成向量
  5. from langchain_community.document_loaders import PyPDFLoader # PDF 文档加载器,把 PDF 解析成带 page_content 和 metadata 的 Document 对象
  6. from langchain_community.vectorstores import Chroma # Chroma 向量数据库,用于存储向量并做相似度检索
  7. from langchain_text_splitters import RecursiveCharacterTextSplitter # 递归字符分块器,按分隔符层级把长文本切成小块
  8. from langchain_openai import ChatOpenAI # DeepSeek API 兼容 OpenAI 接口,用 ChatOpenAI 调用
  9. from langchain_core.prompts import ChatPromptTemplate # 聊天提示词模板,统一管理 system/user 消息格式
  10. from langchain_core.output_parsers import StrOutputParser # 输出解析器,把 LLM 返回的 AIMessage 提取为纯字符串
  11. import os # 标准库,用于读取环境变量(如 DASHSCOPE_API_KEY)
  12. # ========== 第一步:加载文档 ==========
  13. loader = PyPDFLoader('./1.大模型全景认知.pdf')
  14. datas = loader.load()
  15. print(f'原始页数:{len(datas)}')
  16. # ========== 第二步:分块 ==========
  17. text_splitter = RecursiveCharacterTextSplitter(
  18. chunk_size = 200,
  19. chunk_overlap = 40,
  20. separators=['\n\n','\n','。','?','!',' ','']
  21. )
  22. chunks = text_splitter.split_documents(datas)
  23. print(f'分块后的页数:{len(chunks)}')
  24. # ========== 第四步:向量化 + 存入向量库 ==========
  25. embeddings_model = DashScopeEmbeddings(
  26. model='text-embedding-v3',
  27. dashscope_api_key = os.getenv('DASHSCOPE_API_KEY')
  28. )
  29. vectorstore = Chroma.from_documents(
  30. documents=chunks,
  31. embedding=embeddings_model,
  32. collection_metadata={'hnsw:space': 'cosine'},
  33. persist_directory='./chroma_db2',
  34. )
  35. # ========== 第五步:创建检索器 ==========
  36. retriever = vectorstore.as_retriever(
  37. search_type='similarity',
  38. search_kwargs={'k': 3}
  39. )
  40. # ========== 第六步:提问 ==========
  41. query="什么是人工智能,一句话总结"
  42. docs = retriever.invoke(query)
  43. print(docs)
  44. print(f'检索到的文档数:{len(docs)}')
  45. # ========== 第七步:生成回答 ==========
  46. context = '\n\n'.join([d.page_content for d in docs])
  47. prompt = ChatPromptTemplate.from_messages([
  48. ("system", """你是一个专业的知识库助手。请根据以下上下文回答问题。
  49. **规则:**
  50. - 只基于提供的上下文回答,不要编造
  51. - 如果上下文中没有相关信息,直接说「根据现有资料,我找不到这个问题的答案」
  52. - 回答要简洁直接,引用原文时用引号
  53. 参考资料:
  54. {context}"""),
  55. ("human", "{query}")
  56. ])
  57. llm = ChatOpenAI(
  58. model=os.getenv('moduel', 'deepseek-chat'), # .env 中的 moduel 字段
  59. api_key=os.getenv('OPENAI_API_KEY'), # .env 中的 OPENAI_API_KEY
  60. base_url=os.getenv('base_url'), # https://api.deepseek.com
  61. )
  62. chain = prompt | llm | StrOutputParser()
  63. # 执行整条链,获取回答
  64. answer = chain.invoke({"context": context, "query": query})
  65. print(answer)