优化代码,提高解耦
This commit is contained in:
@@ -1,17 +1,11 @@
|
|||||||
import os # 导入环境变量
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
# 让ai说一句话
|
# 让ai说一句话
|
||||||
|
from config import OPENAI_API_KEY
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
|
|
||||||
# 配置deepseek
|
# 配置deepseek
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
# 注:在.env里把OPENAI_API_KEY改成你自己的api-key即可
|
# 注:在.env里把OPENAI_API_KEY改成你自己的api-key即可
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,14 +1,8 @@
|
|||||||
import os # 导入环境变量
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
|
||||||
client = OpenAI(
|
client = OpenAI(
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com")
|
base_url="https://api.deepseek.com")
|
||||||
|
|
||||||
response = client.chat.completions.create(
|
response = client.chat.completions.create(
|
||||||
|
|||||||
@@ -1,12 +1,8 @@
|
|||||||
import os # 导入环境变量
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
|
|
||||||
def create_client():
|
def create_client():
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
return OpenAI(api_key=OPENAI_API_KEY,base_url="https://api.deepseek.com")
|
||||||
return OpenAI(api_key=api_key,base_url="https://api.deepseek.com")
|
|
||||||
|
|
||||||
def chat_loop(agent_client):
|
def chat_loop(agent_client):
|
||||||
messages = [
|
messages = [
|
||||||
|
|||||||
@@ -1,14 +1,10 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
import json
|
import json
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
|
|
||||||
|
|
||||||
def create_client():
|
def create_client():
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
return OpenAI(api_key=OPENAI_API_KEY, base_url="https://api.deepseek.com")
|
||||||
return OpenAI(api_key=api_key, base_url="https://api.deepseek.com")
|
|
||||||
|
|
||||||
def get_weather(location):
|
def get_weather(location):
|
||||||
# 模拟获得天气信息
|
# 模拟获得天气信息
|
||||||
|
|||||||
@@ -1,14 +1,10 @@
|
|||||||
from dotenv import load_dotenv
|
from config import OPENAI_API_KEY
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
import json
|
|
||||||
import os
|
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
|
|
||||||
def create_client():
|
def create_client():
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
return OpenAI(api_key=OPENAI_API_KEY, base_url="https://api.deepseek.com")
|
||||||
return OpenAI(api_key=api_key, base_url="https://api.deepseek.com")
|
|
||||||
|
|
||||||
# 通过api调用获得当前ip与位置
|
# 通过api调用获得当前ip与位置
|
||||||
def get_addr():
|
def get_addr():
|
||||||
|
|||||||
@@ -1,15 +1,10 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
|
|
||||||
# 初始化模型
|
# 初始化模型
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,3 @@
|
|||||||
import os
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
from langchain_core.prompts import ChatPromptTemplate
|
from langchain_core.prompts import ChatPromptTemplate
|
||||||
|
|
||||||
# 定义提示词模板(推荐写法)
|
# 定义提示词模板(推荐写法)
|
||||||
|
|||||||
@@ -1,10 +1,4 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
from langchain_core.prompts import ChatPromptTemplate
|
from langchain_core.prompts import ChatPromptTemplate
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
from langchain_core.output_parsers import StrOutputParser
|
||||||
@@ -18,7 +12,7 @@ prompt = ChatPromptTemplate.from_messages([
|
|||||||
# 初始化模型
|
# 初始化模型
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,18 +1,9 @@
|
|||||||
import os
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
|
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
from langchain_core.output_parsers import StrOutputParser
|
||||||
from langchain_community.chat_message_histories import ChatMessageHistory
|
from langchain_community.chat_message_histories import ChatMessageHistory
|
||||||
from langchain_core.runnables import RunnableWithMessageHistory
|
from langchain_core.runnables import RunnableWithMessageHistory
|
||||||
|
from config import OPENAI_API_KEY
|
||||||
|
|
||||||
prompt = ChatPromptTemplate.from_messages([
|
prompt = ChatPromptTemplate.from_messages([
|
||||||
("system", "你非常可爱,说话末尾会带个喵"),
|
("system", "你非常可爱,说话末尾会带个喵"),
|
||||||
@@ -21,7 +12,7 @@ prompt = ChatPromptTemplate.from_messages([
|
|||||||
])
|
])
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
parser = StrOutputParser()
|
parser = StrOutputParser()
|
||||||
@@ -58,32 +49,3 @@ while 1:
|
|||||||
config={"configurable":{"session_id":session_id}}
|
config={"configurable":{"session_id":session_id}}
|
||||||
)
|
)
|
||||||
print(f'AI:{response}')
|
print(f'AI:{response}')
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,4 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
|
|
||||||
def create_client():
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
return api_key
|
|
||||||
|
|
||||||
|
|
||||||
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
|
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
from langchain_core.output_parsers import StrOutputParser
|
||||||
@@ -44,7 +35,7 @@ def main():
|
|||||||
# 此处集中配置LLM
|
# 此处集中配置LLM
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=create_client(),
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
prompt = """你是‘小智’,一位专业、耐心且记忆力出色的 AI 助手。
|
prompt = """你是‘小智’,一位专业、耐心且记忆力出色的 AI 助手。
|
||||||
|
|||||||
@@ -1,10 +1,3 @@
|
|||||||
import os
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
from langchain_core.tools import tool
|
from langchain_core.tools import tool
|
||||||
|
|
||||||
@tool
|
@tool
|
||||||
|
|||||||
@@ -1,20 +1,15 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
import os
|
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain.agents import AgentExecutor, create_tool_calling_agent
|
# 在LangChain 1.0+版本中,以下俩组件移到了langchain-classic包中
|
||||||
|
from langchain_classic.agents import AgentExecutor
|
||||||
|
from langchain_classic.agents import create_tool_calling_agent
|
||||||
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||||
from langchain_core.tools import tool # 导入 @tool
|
from langchain_core.tools import tool # 导入 @tool
|
||||||
|
|
||||||
# 配置LLM
|
# 配置LLM
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,18 +1,14 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
import sqlite3
|
import sqlite3
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_community.agent_toolkits import create_sql_agent
|
from langchain_community.agent_toolkits import create_sql_agent
|
||||||
from langchain_community.utilities import SQLDatabase # 导入 SQLDatabase
|
from langchain_community.utilities import SQLDatabase # 导入 SQLDatabase
|
||||||
|
import os
|
||||||
|
|
||||||
# 配置llm
|
# 配置llm
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,7 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
import os
|
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
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_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||||
from langchain_core.tools import tool # 导入 @tool
|
from langchain_core.tools import tool # 导入 @tool
|
||||||
from langchain_community.chat_message_histories import ChatMessageHistory
|
from langchain_community.chat_message_histories import ChatMessageHistory
|
||||||
@@ -17,7 +11,7 @@ from langchain_core.runnables import RunnableWithMessageHistory
|
|||||||
# 配置llm
|
# 配置llm
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
# 配置prompt(新增俩占位符 一个为对话历史记录,一个为agent的思考过程)
|
# 配置prompt(新增俩占位符 一个为对话历史记录,一个为agent的思考过程)
|
||||||
|
|||||||
@@ -1,19 +1,13 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
import time
|
import time
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain.globals import set_llm_cache
|
from langchain_core.globals import set_llm_cache
|
||||||
from langchain_community.cache import InMemoryCache
|
from langchain_community.cache import InMemoryCache
|
||||||
|
|
||||||
# 配置llm
|
# 配置llm
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
# 设置全局缓存
|
# 设置全局缓存
|
||||||
|
|||||||
@@ -1,23 +1,18 @@
|
|||||||
import os
|
from config import OPENAI_API_KEY
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
|
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
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_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||||
from langchain_core.tools import tool # 导入 @tool
|
from langchain_core.tools import tool # 导入 @tool
|
||||||
from langchain_community.chat_message_histories import ChatMessageHistory
|
from langchain_community.chat_message_histories import ChatMessageHistory
|
||||||
from langchain_core.runnables import RunnableWithMessageHistory
|
from langchain_core.runnables import RunnableWithMessageHistory
|
||||||
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
|
from langchain_core.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
|
||||||
|
|
||||||
|
|
||||||
# 配置llm
|
# 配置llm
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com",
|
base_url="https://api.deepseek.com",
|
||||||
streaming=True,
|
streaming=True,
|
||||||
callbacks=[StreamingStdOutCallbackHandler()]
|
callbacks=[StreamingStdOutCallbackHandler()]
|
||||||
|
|||||||
@@ -1,12 +1,11 @@
|
|||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
from embeddings import get_embeddings
|
||||||
|
|
||||||
# 第一次运行可能时间较久
|
|
||||||
print('---正在首次加载本地嵌入模型(bge-small-zh-v1.5)...---')
|
# 首次运行可能时间较久 -- 同时运行本文件需要梯子,不然无法加载到本地
|
||||||
|
print('---正在加载本地嵌入模型(bge-small-zh-v1.5)...---')
|
||||||
|
|
||||||
# 理论:有embedding的向量模型
|
# 理论:有embedding的向量模型
|
||||||
embeddings_model = HuggingFaceEmbeddings(
|
embeddings_model = get_embeddings("bge-small-zh-v1.5")
|
||||||
model_name="BAAI/bge-small-zh-v1.5" # 一个中英双语开源模型
|
|
||||||
)
|
|
||||||
print('嵌入模型载入完毕')
|
print('嵌入模型载入完毕')
|
||||||
|
|
||||||
# 演示:将文本转换为向量
|
# 演示:将文本转换为向量
|
||||||
|
|||||||
@@ -1,10 +1,9 @@
|
|||||||
# 仅负责: 切片 -> 向量化 -> 构建索引 -> 保存到磁盘
|
# 仅负责: 切片 -> 向量化 -> 构建索引 -> 保存到磁盘
|
||||||
# 运行一次即可,无需每次检索都运行
|
# 运行一次即可,无需每次检索都运行
|
||||||
|
from embeddings import get_embeddings
|
||||||
import os
|
import os
|
||||||
from langchain_community.document_loaders import TextLoader
|
from langchain_community.document_loaders import TextLoader
|
||||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
from langchain_community.vectorstores import FAISS
|
from langchain_community.vectorstores import FAISS
|
||||||
|
|
||||||
# 准备知识库内容
|
# 准备知识库内容
|
||||||
@@ -48,7 +47,7 @@ splits = text_splitter.split_documents(docs) # 运行切分器
|
|||||||
print(f'p1完成,文档已切分成{len(splits)}个片段\n')
|
print(f'p1完成,文档已切分成{len(splits)}个片段\n')
|
||||||
|
|
||||||
# 2. 向量化(Embedding)
|
# 2. 向量化(Embedding)
|
||||||
embeddings_model = HuggingFaceEmbeddings(model_name="BAAI/bge-small-zh-v1.5") # 载入向量化模型
|
embeddings_model = get_embeddings() # 载入向量化模型
|
||||||
print(f'p2完成,Embedding模型已准备\n') #
|
print(f'p2完成,Embedding模型已准备\n') #
|
||||||
|
|
||||||
# 3. 存储(Store)
|
# 3. 存储(Store)
|
||||||
@@ -61,17 +60,3 @@ print(f'p3完成,向量数据库{db}已构建')
|
|||||||
os.remove("knowledge_base.txt")
|
os.remove("knowledge_base.txt")
|
||||||
|
|
||||||
print('---所有阶段已经完成!---')
|
print('---所有阶段已经完成!---')
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,12 +1,8 @@
|
|||||||
|
from embeddings import get_embeddings
|
||||||
|
from config import OPENAI_API_KEY
|
||||||
import os
|
import os
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
from langchain_community.document_loaders import TextLoader
|
from langchain_community.document_loaders import TextLoader
|
||||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
from langchain_community.vectorstores import FAISS
|
from langchain_community.vectorstores import FAISS
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.prompts import ChatPromptTemplate
|
from langchain_core.prompts import ChatPromptTemplate
|
||||||
@@ -43,8 +39,8 @@ loader = TextLoader('knowledge_base.txt',encoding='utf8')
|
|||||||
docs = loader.load()
|
docs = loader.load()
|
||||||
text_splitter = RecursiveCharacterTextSplitter(chunk_size=250,chunk_overlap=40)
|
text_splitter = RecursiveCharacterTextSplitter(chunk_size=250,chunk_overlap=40)
|
||||||
splits = text_splitter.split_documents(docs)
|
splits = text_splitter.split_documents(docs)
|
||||||
embedding_model = HuggingFaceEmbeddings(model_name="BAAI/bge-small-zh-v1.5")
|
embedding_model = get_embeddings("bge-small-zh-v1.5")
|
||||||
db = FAISS.from_documents(splits,embedding_model)
|
db = FAISS.from_documents(splits,embedding_model) # 在内存中构建向量索引,但不持久化到本地文件
|
||||||
print('---模块A(Indexing)完成---\n')
|
print('---模块A(Indexing)完成---\n')
|
||||||
|
|
||||||
|
|
||||||
@@ -75,7 +71,7 @@ prompt = ChatPromptTemplate.from_messages([
|
|||||||
# 3. G (Generation - 生成)
|
# 3. G (Generation - 生成)
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -111,15 +107,3 @@ print(f'回答:{response}\n')
|
|||||||
|
|
||||||
# 清理临时文件
|
# 清理临时文件
|
||||||
os.remove("knowledge_base.txt")
|
os.remove("knowledge_base.txt")
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,17 +1,14 @@
|
|||||||
# pip install chroma
|
|
||||||
# pip install -U langchain-chroma
|
|
||||||
|
|
||||||
import os
|
import os
|
||||||
from langchain_community.document_loaders import TextLoader
|
from langchain_community.document_loaders import TextLoader
|
||||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
from langchain_chroma import Chroma
|
from langchain_chroma import Chroma
|
||||||
|
from embeddings import get_embeddings
|
||||||
|
|
||||||
|
|
||||||
knowledge_base_file = "war_and_peace.txt"
|
knowledge_base_file = "war_and_peace.txt"
|
||||||
# 持久化目录: Chroma会把所有数据(向量+文本+元数据)都存到这个文件夹
|
# 持久化目录: Chroma会把所有数据(向量+文本+元数据)都存到这个文件夹
|
||||||
persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
|
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_size = 500
|
||||||
chunk_overlap = 75
|
chunk_overlap = 75
|
||||||
|
|
||||||
@@ -42,9 +39,10 @@ text_splitter = RecursiveCharacterTextSplitter(
|
|||||||
splits = text_splitter.split_documents(docs)
|
splits = text_splitter.split_documents(docs)
|
||||||
print('分割完成...\n')
|
print('分割完成...\n')
|
||||||
# 3. 向量化 -- 第一次运行会下载模型,预计耗时2分钟
|
# 3. 向量化 -- 第一次运行会下载模型,预计耗时2分钟
|
||||||
embedding_model = HuggingFaceEmbeddings(
|
print(f'正在加载/下载模型{model_name_str}...')
|
||||||
model_name=embedding_model,
|
embedding_model = get_embeddings(
|
||||||
model_kwargs={'device':'cpu'}, # 强制模型在cpu上运行
|
model_name=model_name_str,
|
||||||
|
device='cpu', # 强制模型在cpu上运行
|
||||||
encode_kwargs={'batch_size':64} # 每次处理64个文本片段
|
encode_kwargs={'batch_size':64} # 每次处理64个文本片段
|
||||||
)
|
)
|
||||||
print('Embedding模型加载完成...\n')
|
print('Embedding模型加载完成...\n')
|
||||||
@@ -64,9 +62,3 @@ for i in range(0, len(splits), batch_size):
|
|||||||
|
|
||||||
|
|
||||||
print(f'✅ 索引构建完毕,共 {len(splits)} 条,已保存到 {persist_directory}')
|
print(f'✅ 索引构建完毕,共 {len(splits)} 条,已保存到 {persist_directory}')
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -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_chroma import Chroma
|
||||||
from langchain_openai import ChatOpenAI
|
|
||||||
from langchain_core.prompts import ChatPromptTemplate
|
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
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'
|
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):
|
if not os.path.exists(Persist_directory):
|
||||||
print(f"错误: 知识库文件 {Persist_directory} 未找到。")
|
print(f"错误: 知识库文件 {Persist_directory} 未找到。")
|
||||||
print("请先运行'build_index.py'生成向量数据库,再运行该文件")
|
print("请先运行'00_build_index.py'生成向量数据库,再运行该文件")
|
||||||
exit()
|
exit()
|
||||||
|
|
||||||
print('---加载本地向量数据库---')
|
print('---加载本地向量数据库---')
|
||||||
|
|
||||||
# 模块A:链接本地Chroma向量数据库
|
# 模块A:链接本地Chroma向量数据库
|
||||||
# 1. 加载 Embedding 模型
|
# 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
|
# 2. 从本地目录加载Chroma DB
|
||||||
db = Chroma(
|
db = Chroma(
|
||||||
persist_directory=Persist_directory,
|
persist_directory=Persist_directory,
|
||||||
embedding_function=embedding_model
|
embedding_function=embeddings_model
|
||||||
)
|
)
|
||||||
print(f'Chroma数据库已从本地加载(共{db._collection.count()}条)\n')
|
print(f'Chroma数据库已从本地加载(共{db._collection.count()}条)\n')
|
||||||
|
|
||||||
# 模块B:R-A-G Flow
|
# 模块B:R-A-G Flow
|
||||||
# 1. R-检索
|
# 1. R-检索
|
||||||
retriever = db.as_retriever(search_kwargs={"k": 3}) # 召回3条相关数据
|
retriever = db.as_retriever(search_kwargs={"k": 5}) # 召回5条相关数据
|
||||||
|
|
||||||
# 2. A-增强
|
# 2. A-增强
|
||||||
sys_prompt = """
|
sys_prompt = """
|
||||||
@@ -58,7 +54,7 @@ prompt = ChatPromptTemplate.from_messages([
|
|||||||
# 3. G-生成
|
# 3. G-生成
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -80,5 +76,3 @@ question = '莫斯科大火发生在小说的哪一部分?有哪些角色亲
|
|||||||
response = rag_chain.invoke(question)
|
response = rag_chain.invoke(question)
|
||||||
print(f'提问:{question}')
|
print(f'提问:{question}')
|
||||||
print(f'回答:{response}')
|
print(f'回答:{response}')
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,33 +1,31 @@
|
|||||||
import os
|
import os
|
||||||
from dotenv import load_dotenv
|
from config import OPENAI_API_KEY
|
||||||
|
from embeddings import get_embeddings
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
from langchain_chroma import Chroma
|
from langchain_chroma import Chroma
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.prompts import ChatPromptTemplate
|
from langchain_core.prompts import ChatPromptTemplate
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
from langchain_core.output_parsers import StrOutputParser
|
||||||
from langchain_core.runnables import RunnablePassthrough
|
from langchain_core.runnables import RunnablePassthrough
|
||||||
|
from langchain_classic.retrievers import ContextualCompressionRetriever
|
||||||
# --- Reranker (02) 新增的 import ---
|
from langchain_classic.retrievers.document_compressors import CrossEncoderReranker
|
||||||
from langchain.retrievers import ContextualCompressionRetriever
|
|
||||||
from langchain_community.cross_encoders import HuggingFaceCrossEncoder
|
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'
|
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):
|
if not os.path.exists(Persist_directory):
|
||||||
print(f"错误: 知识库文件 {Persist_directory} 未找到。")
|
print(f"错误: 知识库文件 {Persist_directory} 未找到。")
|
||||||
print("请先运行'build_index.py'生成向量数据库,再运行该文件")
|
print("请先运行'00_build_index.py'生成向量数据库,再运行该文件")
|
||||||
exit()
|
exit()
|
||||||
|
|
||||||
print('---加载本地向量数据库---\n')
|
print('---加载本地向量数据库---\n')
|
||||||
|
|
||||||
# 1. 加载 Embedding 模型
|
# 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
|
# 2. 加载 Chroma db
|
||||||
db = Chroma(
|
db = Chroma(
|
||||||
@@ -40,11 +38,11 @@ print('---Chroma数据库已加载---\n')
|
|||||||
# --- 模块 B (R-A-G Flow) ---
|
# --- 模块 B (R-A-G Flow) ---
|
||||||
# 1. R-检索--强化版
|
# 1. R-检索--强化版
|
||||||
# 1.1 基础检索器(Base Retriever) - '粗召回'
|
# 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 (重排器) - "精排序" -- 首次运行需要耗时下载
|
# 1.2 Reranker (重排器) - "精排序" -- 首次运行需要耗时下载
|
||||||
print('正在加载Reranker模型 (bge-reranker-base)...')
|
print('正在加载 Reranker模型 (bge-reranker-base)...')
|
||||||
encoder = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-base") # 加载Ranker模型
|
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 创建管道封装器
|
# 1.3 创建管道封装器
|
||||||
compression_retriever = ContextualCompressionRetriever(
|
compression_retriever = ContextualCompressionRetriever(
|
||||||
base_retriever=base_retriever, # 用Chroma做 海选
|
base_retriever=base_retriever, # 用Chroma做 海选
|
||||||
@@ -71,7 +69,7 @@ prompt = ChatPromptTemplate.from_messages([
|
|||||||
# 3. G-生成
|
# 3. G-生成
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -94,17 +92,3 @@ question = '皮埃尔是共济会成员吗?他在其中扮演什么角色?'
|
|||||||
response = rag_chain.invoke(question)
|
response = rag_chain.invoke(question)
|
||||||
print(f'提问:{question}')
|
print(f'提问:{question}')
|
||||||
print(f'回答:{response}')
|
print(f'回答:{response}')
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,25 +1,21 @@
|
|||||||
import os
|
import os
|
||||||
from dotenv import load_dotenv
|
from config import OPENAI_API_KEY
|
||||||
|
from embeddings import get_embeddings
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
from langchain_chroma import Chroma
|
from langchain_chroma import Chroma
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.prompts import ChatPromptTemplate
|
from langchain_core.prompts import ChatPromptTemplate
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
from langchain_core.output_parsers import StrOutputParser
|
||||||
from langchain_core.runnables import RunnablePassthrough
|
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_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_core.tools import tool
|
||||||
|
|
||||||
|
|
||||||
# 全局 LLM (供Agent和Rag共用)
|
# 全局 LLM (供Agent和Rag共用)
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -27,23 +23,26 @@ llm = ChatOpenAI(
|
|||||||
def build_rag_chain(llm_instance):
|
def build_rag_chain(llm_instance):
|
||||||
print('---正在构建RAG链条...---\n')
|
print('---正在构建RAG链条...---\n')
|
||||||
|
|
||||||
Persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
|
persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
|
||||||
Embedding_model = 'BAAI/bge-small-en-v1.5'
|
embedding_model_name = 'BAAI/bge-small-en-v1.5'
|
||||||
Encoder_model = "BAAI/bge-reranker-base"
|
encoder_model_name = "BAAI/bge-reranker-base"
|
||||||
|
|
||||||
if not os.path.exists(Persist_directory):
|
if not os.path.exists(persist_directory):
|
||||||
raise FileNotFoundError(f'索引目录{Persist_directory}未找到,请先运行 build_index.py')
|
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(
|
db = Chroma(
|
||||||
persist_directory=Persist_directory,
|
persist_directory=persist_directory,
|
||||||
embedding_function=embeddings_model
|
embedding_function=embeddings_model
|
||||||
)
|
)
|
||||||
|
|
||||||
# 1. R-检索--强化版
|
# 1. R-检索--强化版
|
||||||
base_retriever = db.as_retriever(search_kwargs={"k":5})
|
base_retriever = db.as_retriever(search_kwargs={"k":50})
|
||||||
encoder = HuggingFaceCrossEncoder(model_name=Encoder_model)
|
|
||||||
reranker = CrossEncoderReranker(model=encoder,top_n=2)
|
print(f'正在加载 Reranker模型:{encoder_model_name}...')
|
||||||
|
encoder = HuggingFaceCrossEncoder(model_name=encoder_model_name)
|
||||||
|
reranker = CrossEncoderReranker(model=encoder,top_n=6)
|
||||||
compression_retriever=ContextualCompressionRetriever(
|
compression_retriever=ContextualCompressionRetriever(
|
||||||
base_retriever=base_retriever,
|
base_retriever=base_retriever,
|
||||||
base_compressor=reranker
|
base_compressor=reranker
|
||||||
@@ -59,6 +58,7 @@ def build_rag_chain(llm_instance):
|
|||||||
[上下文]: {context}
|
[上下文]: {context}
|
||||||
[问题]: {question}
|
[问题]: {question}
|
||||||
"""
|
"""
|
||||||
|
|
||||||
prompt = ChatPromptTemplate.from_messages([
|
prompt = ChatPromptTemplate.from_messages([
|
||||||
('system',sys_prompt),
|
('system',sys_prompt),
|
||||||
('human','{question}')
|
('human','{question}')
|
||||||
@@ -106,31 +106,3 @@ if __name__ == '__main__':
|
|||||||
res = search_war_and_peace.invoke(question)
|
res = search_war_and_peace.invoke(question)
|
||||||
print(f'问题:{question}')
|
print(f'问题:{question}')
|
||||||
print(f'回答:{res}')
|
print(f'回答:{res}')
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,21 +1,17 @@
|
|||||||
import os
|
import os
|
||||||
from dotenv import load_dotenv
|
from config import OPENAI_API_KEY
|
||||||
|
from embeddings import get_embeddings
|
||||||
load_dotenv()
|
|
||||||
api_key = os.getenv("OPENAI_API_KEY")
|
|
||||||
|
|
||||||
from langchain_huggingface import HuggingFaceEmbeddings
|
|
||||||
from langchain_chroma import Chroma
|
from langchain_chroma import Chroma
|
||||||
from langchain_openai import ChatOpenAI
|
from langchain_openai import ChatOpenAI
|
||||||
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
|
from langchain_core.prompts import ChatPromptTemplate,MessagesPlaceholder
|
||||||
from langchain_core.output_parsers import StrOutputParser
|
from langchain_core.output_parsers import StrOutputParser
|
||||||
from langchain_core.runnables import RunnablePassthrough
|
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_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_core.tools import tool
|
||||||
|
from langchain_classic.agents import AgentExecutor
|
||||||
from langchain.agents import AgentExecutor, create_tool_calling_agent
|
from langchain_classic.agents import create_tool_calling_agent
|
||||||
from langchain_community.chat_message_histories import ChatMessageHistory
|
from langchain_community.chat_message_histories import ChatMessageHistory
|
||||||
from langchain_core.runnables import RunnableWithMessageHistory
|
from langchain_core.runnables import RunnableWithMessageHistory
|
||||||
|
|
||||||
@@ -25,23 +21,25 @@ from langchain_core.runnables import RunnableWithMessageHistory
|
|||||||
def build_rag_chain(llm_instance):
|
def build_rag_chain(llm_instance):
|
||||||
print('---正在构建RAG链条---')
|
print('---正在构建RAG链条---')
|
||||||
|
|
||||||
Persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
|
persist_directory = './chroma_db_war_and_peace_bge_small_en_v1.5'
|
||||||
Embedding_model = 'BAAI/bge-small-en-v1.5'
|
embedding_model_name = 'BAAI/bge-small-en-v1.5'
|
||||||
Encoder_model = "BAAI/bge-reranker-base"
|
encoder_model_name = "BAAI/bge-reranker-base"
|
||||||
|
|
||||||
if not os.path.exists(Persist_directory):
|
if not os.path.exists(persist_directory):
|
||||||
raise FileNotFoundError(f'索引目录{Persist_directory}未找到,请先运行 build_index.py')
|
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(
|
db = Chroma(
|
||||||
persist_directory=Persist_directory,
|
persist_directory=persist_directory,
|
||||||
embedding_function=embedding_model
|
embedding_function=embeddings_model
|
||||||
)
|
)
|
||||||
# R
|
# R
|
||||||
base_retriever = db.as_retriever(search_kwargs={'k':5})
|
print(f'正在加载 Reranker模型:{encoder_model_name}...')
|
||||||
encoder = HuggingFaceCrossEncoder(model_name=Encoder_model)
|
base_retriever = db.as_retriever(search_kwargs={'k':50})
|
||||||
reranker = CrossEncoderReranker(model=encoder,top_n=2)
|
encoder = HuggingFaceCrossEncoder(model_name=encoder_model_name)
|
||||||
|
reranker = CrossEncoderReranker(model=encoder,top_n=6)
|
||||||
compression_retriever = ContextualCompressionRetriever(
|
compression_retriever = ContextualCompressionRetriever(
|
||||||
base_retriever=base_retriever,
|
base_retriever=base_retriever,
|
||||||
base_compressor=reranker
|
base_compressor=reranker
|
||||||
@@ -80,7 +78,7 @@ def create_agent_with_memory():
|
|||||||
# LLm
|
# LLm
|
||||||
llm = ChatOpenAI(
|
llm = ChatOpenAI(
|
||||||
model="deepseek-chat",
|
model="deepseek-chat",
|
||||||
api_key=api_key,
|
api_key=OPENAI_API_KEY,
|
||||||
base_url="https://api.deepseek.com"
|
base_url="https://api.deepseek.com"
|
||||||
)
|
)
|
||||||
# Prompt
|
# Prompt
|
||||||
|
|||||||
+14
-12
@@ -1,12 +1,14 @@
|
|||||||
python-dotenv~=1.0.0
|
huggingface_hub==0.36.0
|
||||||
openai>=1.40.0,<2.0.0
|
langchain==1.0.8
|
||||||
requests~=2.32.4
|
langchain_chroma==1.0.0
|
||||||
langchain~=0.2.16
|
langchain_classic==1.0.0
|
||||||
langchain-community~=0.2.16
|
langchain_community==0.4.1
|
||||||
langchain-openai~=0.1.15
|
langchain_core==1.0.6
|
||||||
langchain-core~=0.2.38
|
langchain_huggingface==1.0.1
|
||||||
langgraph~=0.2.0
|
langchain_openai==1.0.3
|
||||||
langchain-huggingface==0.0.3
|
langchain_text_splitters==1.0.0
|
||||||
langchain-chroma==0.1.1
|
langgraph==1.0.3
|
||||||
chromadb==0.4.22
|
langsmith==0.4.43
|
||||||
sentence-transformers>=2.2.0
|
openai==2.8.1
|
||||||
|
python-dotenv==1.2.1
|
||||||
|
Requests==2.32.5
|
||||||
|
|||||||
Reference in New Issue
Block a user