Ver Fonte

添加 rag_chain.py 到 01_work

Daphne chen há 2 meses atrás
commit
060098e2b1
1 ficheiros alterados com 81 adições e 0 exclusões
  1. 81 0
      01_work/rag_chain.py

+ 81 - 0
01_work/rag_chain.py

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