优化代码,提高解耦

This commit is contained in:
Annyfee
2025-11-19 23:44:14 +08:00
parent 8265c2a7a9
commit a70517f6d8
25 changed files with 143 additions and 355 deletions
@@ -1,17 +1,14 @@
# pip install chroma
# pip install -U langchain-chroma
import os
from langchain_community.document_loaders import TextLoader
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_chroma import Chroma
from embeddings import get_embeddings
knowledge_base_file = "war_and_peace.txt"
# 持久化目录: Chroma会把所有数据(向量+文本+元数据)都存到这个文件夹
persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
embedding_model = 'BAAI/bge-small-en-v1.5' # 如果愿意等待,可以换成模型"BAAI/bge-m3",效果更好更适合长文,但下载时间也更久(2.2G)
model_name_str = 'BAAI/bge-small-en-v1.5' # 如果愿意等待,可以换成模型"BAAI/bge-m3",效果更好更适合长文,但下载时间也更久(2.2G)
chunk_size = 500
chunk_overlap = 75
@@ -42,9 +39,10 @@ text_splitter = RecursiveCharacterTextSplitter(
splits = text_splitter.split_documents(docs)
print('分割完成...\n')
# 3. 向量化 -- 第一次运行会下载模型,预计耗时2分钟
embedding_model = HuggingFaceEmbeddings(
model_name=embedding_model,
model_kwargs={'device':'cpu'}, # 强制模型在cpu上运行
print(f'正在加载/下载模型{model_name_str}...')
embedding_model = get_embeddings(
model_name=model_name_str,
device='cpu', # 强制模型在cpu上运行
encode_kwargs={'batch_size':64} # 每次处理64个文本片段
)
print('Embedding模型加载完成...\n')
@@ -63,10 +61,4 @@ for i in range(0, len(splits), batch_size):
print(f"已插入 {min(i + batch_size, len(splits))} / {len(splits)}")
print(f'✅ 索引构建完毕,共 {len(splits)} 条,已保存到 {persist_directory}')
print(f'✅ 索引构建完毕,共 {len(splits)} 条,已保存到 {persist_directory}')
+17 -23
View File
@@ -1,45 +1,41 @@
# pip install --upgrade langchain-openai
# pip install --upgrade langchain-huggingface langchain-core langchain-community
# pip install --upgrade langchain-core langchain-community
import os
from dotenv import load_dotenv
from langchain_core.runnables import RunnablePassthrough
load_dotenv()
api_key = os.getenv("OPENAI_API_KEY")
import os
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_chroma import Chroma
from langchain_openai import ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_openai import ChatOpenAI
from embeddings import get_embeddings
from config import OPENAI_API_KEY
import os
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.runnables import RunnablePassthrough
Persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
Embedding_model = 'BAAI/bge-small-en-v1.5'
model_name_str = 'BAAI/bge-small-en-v1.5'
if not os.path.exists(Persist_directory):
print(f"错误: 知识库文件 {Persist_directory} 未找到。")
print("请先运行'build_index.py'生成向量数据库,再运行该文件")
print("请先运行'00_build_index.py'生成向量数据库,再运行该文件")
exit()
print('---加载本地向量数据库---')
# 模块A:链接本地Chroma向量数据库
# 1. 加载 Embedding 模型
embedding_model = HuggingFaceEmbeddings(model_name=Embedding_model)
print(f'正在加载/下载模型{model_name_str}...')
embeddings_model = get_embeddings(
model_name=model_name_str,
device='cpu'
)
# 2. 从本地目录加载Chroma DB
db = Chroma(
persist_directory=Persist_directory,
embedding_function=embedding_model
embedding_function=embeddings_model
)
print(f'Chroma数据库已从本地加载(共{db._collection.count()}条)\n')
# 模块B:R-A-G Flow
# 1. R-检索
retriever = db.as_retriever(search_kwargs={"k": 3}) # 召回3条相关数据
retriever = db.as_retriever(search_kwargs={"k": 5}) # 召回5条相关数据
# 2. A-增强
sys_prompt = """
@@ -58,7 +54,7 @@ prompt = ChatPromptTemplate.from_messages([
# 3. G-生成
llm = ChatOpenAI(
model="deepseek-chat",
api_key=api_key,
api_key=OPENAI_API_KEY,
base_url="https://api.deepseek.com"
)
@@ -79,6 +75,4 @@ print('---正在运行RAG链条---')
question = '莫斯科大火发生在小说的哪一部分?有哪些角色亲历了这场灾难?'
response = rag_chain.invoke(question)
print(f'提问:{question}')
print(f'回答:{response}')
print(f'回答:{response}')
+15 -31
View File
@@ -1,33 +1,31 @@
import os
from dotenv import load_dotenv
load_dotenv()
api_key = os.getenv("OPENAI_API_KEY")
from langchain_huggingface import HuggingFaceEmbeddings
from config import OPENAI_API_KEY
from embeddings import get_embeddings
from langchain_chroma import Chroma
from langchain_openai import ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough
# --- Reranker (02) 新增的 import ---
from langchain.retrievers import ContextualCompressionRetriever
from langchain_classic.retrievers import ContextualCompressionRetriever
from langchain_classic.retrievers.document_compressors import CrossEncoderReranker
from langchain_community.cross_encoders import HuggingFaceCrossEncoder
from langchain.retrievers.document_compressors import CrossEncoderReranker
Persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
Embedding_model = 'BAAI/bge-small-en-v1.5'
model_name_str = 'BAAI/bge-small-en-v1.5'
if not os.path.exists(Persist_directory):
print(f"错误: 知识库文件 {Persist_directory} 未找到。")
print("请先运行'build_index.py'生成向量数据库,再运行该文件")
print("请先运行'00_build_index.py'生成向量数据库,再运行该文件")
exit()
print('---加载本地向量数据库---\n')
# 1. 加载 Embedding 模型
embeddings_model = HuggingFaceEmbeddings(model_name=Embedding_model)
print(f'正在加载/下载模型{model_name_str}...')
embeddings_model = get_embeddings(
model_name=model_name_str,
device='cpu'
)
# 2. 加载 Chroma db
db = Chroma(
@@ -40,11 +38,11 @@ print('---Chroma数据库已加载---\n')
# --- 模块 B (R-A-G Flow) ---
# 1. R-检索--强化版
# 1.1 基础检索器(Base Retriever) - '粗召回'
base_retriever = db.as_retriever(search_kwargs={"k":5}) # K调大到5
base_retriever = db.as_retriever(search_kwargs={"k":50}) # K调大到60
# 1.2 Reranker (重排器) - "精排序" -- 首次运行需要耗时下载
print('正在加载Reranker模型 (bge-reranker-base)...')
print('正在加载 Reranker模型 (bge-reranker-base)...')
encoder = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-base") # 加载Ranker模型
reranker = CrossEncoderReranker(model=encoder,top_n=2) # 对检索结果进行精排
reranker = CrossEncoderReranker(model=encoder,top_n=6) # 对检索结果进行精排
# 1.3 创建管道封装器
compression_retriever = ContextualCompressionRetriever(
base_retriever=base_retriever, # 用Chroma做 海选
@@ -71,7 +69,7 @@ prompt = ChatPromptTemplate.from_messages([
# 3. G-生成
llm = ChatOpenAI(
model="deepseek-chat",
api_key=api_key,
api_key=OPENAI_API_KEY,
base_url="https://api.deepseek.com"
)
@@ -94,17 +92,3 @@ question = '皮埃尔是共济会成员吗?他在其中扮演什么角色?'
response = rag_chain.invoke(question)
print(f'提问:{question}')
print(f'回答:{response}')
+20 -48
View File
@@ -1,25 +1,21 @@
import os
from dotenv import load_dotenv
load_dotenv()
api_key = os.getenv("OPENAI_API_KEY")
from langchain_huggingface import HuggingFaceEmbeddings
from config import OPENAI_API_KEY
from embeddings import get_embeddings
from langchain_chroma import Chroma
from langchain_openai import ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough
from langchain.retrievers import ContextualCompressionRetriever
from langchain_classic.retrievers import ContextualCompressionRetriever
from langchain_community.cross_encoders import HuggingFaceCrossEncoder
from langchain.retrievers.document_compressors import CrossEncoderReranker
from langchain_classic.retrievers.document_compressors import CrossEncoderReranker
from langchain_core.tools import tool
# 全局 LLM (供Agent和Rag共用)
llm = ChatOpenAI(
model="deepseek-chat",
api_key=api_key,
api_key=OPENAI_API_KEY,
base_url="https://api.deepseek.com"
)
@@ -27,23 +23,26 @@ llm = ChatOpenAI(
def build_rag_chain(llm_instance):
print('---正在构建RAG链条...---\n')
Persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
Embedding_model = 'BAAI/bge-small-en-v1.5'
Encoder_model = "BAAI/bge-reranker-base"
persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
embedding_model_name = 'BAAI/bge-small-en-v1.5'
encoder_model_name = "BAAI/bge-reranker-base"
if not os.path.exists(Persist_directory):
raise FileNotFoundError(f'索引目录{Persist_directory}未找到,请先运行 build_index.py')
if not os.path.exists(persist_directory):
raise FileNotFoundError(f'索引目录{persist_directory}未找到,请先运行 build_index.py')
embeddings_model = HuggingFaceEmbeddings(model_name=Embedding_model)
print(f'正在加载/下载 Embedding模型:{embedding_model_name}')
embeddings_model = get_embeddings(model_name=embedding_model_name,device='cpu')
db = Chroma(
persist_directory=Persist_directory,
persist_directory=persist_directory,
embedding_function=embeddings_model
)
# 1. R-检索--强化版
base_retriever = db.as_retriever(search_kwargs={"k":5})
encoder = HuggingFaceCrossEncoder(model_name=Encoder_model)
reranker = CrossEncoderReranker(model=encoder,top_n=2)
base_retriever = db.as_retriever(search_kwargs={"k":50})
print(f'正在加载 Reranker模型:{encoder_model_name}...')
encoder = HuggingFaceCrossEncoder(model_name=encoder_model_name)
reranker = CrossEncoderReranker(model=encoder,top_n=6)
compression_retriever=ContextualCompressionRetriever(
base_retriever=base_retriever,
base_compressor=reranker
@@ -59,6 +58,7 @@ def build_rag_chain(llm_instance):
[上下文]: {context}
[问题]: {question}
"""
prompt = ChatPromptTemplate.from_messages([
('system',sys_prompt),
('human','{question}')
@@ -105,32 +105,4 @@ if __name__ == '__main__':
question = "皮埃尔是共济会成员吗?他在其中扮演什么角色?"
res = search_war_and_peace.invoke(question)
print(f'问题:{question}')
print(f'回答:{res}')
print(f'回答:{res}')
+20 -22
View File
@@ -1,21 +1,17 @@
import os
from dotenv import load_dotenv
load_dotenv()
api_key = os.getenv("OPENAI_API_KEY")
from langchain_huggingface import HuggingFaceEmbeddings
from config import OPENAI_API_KEY
from embeddings import get_embeddings
from langchain_chroma import Chroma
from langchain_openai import ChatOpenAI
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnablePassthrough
from langchain.retrievers import ContextualCompressionRetriever
from langchain_classic.retrievers import ContextualCompressionRetriever
from langchain_community.cross_encoders import HuggingFaceCrossEncoder
from langchain.retrievers.document_compressors import CrossEncoderReranker
from langchain_classic.retrievers.document_compressors import CrossEncoderReranker
from langchain_core.tools import tool
from langchain.agents import AgentExecutor, create_tool_calling_agent
from langchain_classic.agents import AgentExecutor
from langchain_classic.agents import create_tool_calling_agent
from langchain_community.chat_message_histories import ChatMessageHistory
from langchain_core.runnables import RunnableWithMessageHistory
@@ -25,23 +21,25 @@ from langchain_core.runnables import RunnableWithMessageHistory
def build_rag_chain(llm_instance):
print('---正在构建RAG链条---')
Persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
Embedding_model = 'BAAI/bge-small-en-v1.5'
Encoder_model = "BAAI/bge-reranker-base"
persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
embedding_model_name = 'BAAI/bge-small-en-v1.5'
encoder_model_name = "BAAI/bge-reranker-base"
if not os.path.exists(Persist_directory):
raise FileNotFoundError(f'索引目录{Persist_directory}未找到,请先运行 build_index.py')
if not os.path.exists(persist_directory):
raise FileNotFoundError(f'索引目录{persist_directory}未找到,请先运行 build_index.py')
# 链接向量数据库
embedding_model = HuggingFaceEmbeddings(model_name=Embedding_model)
print(f'正在加载/下载 Embedding模型:{embedding_model_name}')
embeddings_model = get_embeddings(model_name=embedding_model_name,device='cpu')
db = Chroma(
persist_directory=Persist_directory,
embedding_function=embedding_model
persist_directory=persist_directory,
embedding_function=embeddings_model
)
# R
base_retriever = db.as_retriever(search_kwargs={'k':5})
encoder = HuggingFaceCrossEncoder(model_name=Encoder_model)
reranker = CrossEncoderReranker(model=encoder,top_n=2)
print(f'正在加载 Reranker模型:{encoder_model_name}...')
base_retriever = db.as_retriever(search_kwargs={'k':50})
encoder = HuggingFaceCrossEncoder(model_name=encoder_model_name)
reranker = CrossEncoderReranker(model=encoder,top_n=6)
compression_retriever = ContextualCompressionRetriever(
base_retriever=base_retriever,
base_compressor=reranker
@@ -80,7 +78,7 @@ def create_agent_with_memory():
# LLm
llm = ChatOpenAI(
model="deepseek-chat",
api_key=api_key,
api_key=OPENAI_API_KEY,
base_url="https://api.deepseek.com"
)
# Prompt