Merge pull request #149 from xerrors/db_tool_update

优化知识库查询以及工具定义方法
This commit is contained in:
Wenjie Zhang 2025-04-11 12:25:19 +08:00 committed by GitHub
commit 1d20db6d15
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 103 additions and 23 deletions

View File

@ -66,19 +66,25 @@ def regist_tool(
class KnowledgeRetrieverModel(BaseModel): class KnowledgeRetrieverModel(BaseModel):
query: str = Field(description="The query to get knowledge graph.") query: str = Field(description="查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。")
def get_all_tools(): def get_all_tools():
"""获取所有工具""" """获取所有工具"""
tools = _TOOLS_REGISTRY.copy() tools = _TOOLS_REGISTRY.copy()
# 获取所有知识库
for db_Id, retrieve_info in knowledge_base.get_retrievers().items(): for db_Id, retrieve_info in knowledge_base.get_retrievers().items():
name = f"retrieve_{retrieve_info['name']}" name = f"retrieve_{retrieve_info['name']}"
description = (
f"使用 {retrieve_info['name']} 知识库进行检索。\n"
f"下面是这个知识库的描述:\n{retrieve_info['description']}"
)
tools[name] = StructuredTool.from_function( tools[name] = StructuredTool.from_function(
retrieve_info["retriever"], retrieve_info["retriever"],
name=name, name=name,
description=retrieve_info["description"], description=description,
args_schema=KnowledgeRetrieverModel) args_schema=KnowledgeRetrieverModel)
return tools return tools

View File

@ -11,7 +11,6 @@ from src.utils import logger, hashstr
from src.core.indexing import chunk, read_text from src.core.indexing import chunk, read_text
from src.core.kb_db_manager import kb_db_manager from src.core.kb_db_manager import kb_db_manager
class KnowledgeBase: class KnowledgeBase:
def __init__(self) -> None: def __init__(self) -> None:
@ -250,7 +249,7 @@ class KnowledgeBase:
def add_files(self, db_id, files, params=None): def add_files(self, db_id, files, params=None):
db = self.get_kb_by_id(db_id) db = self.get_kb_by_id(db_id)
if db["embed_model"] != self.embed_model.embed_model_fullname: if not self.check_embed_model(db_id):
logger.error(f"Embed model not match, {db['embed_model']} != {self.embed_model.embed_model_fullname}") logger.error(f"Embed model not match, {db['embed_model']} != {self.embed_model.embed_model_fullname}")
return {"message": f"Embed model not match, cur: {self.embed_model.embed_model_fullname}, req: {db['embed_model']}", "status": "failed"} return {"message": f"Embed model not match, cur: {self.embed_model.embed_model_fullname}, req: {db['embed_model']}", "status": "failed"}
@ -312,7 +311,6 @@ class KnowledgeBase:
################################### ###################################
def query(self, query, db_id, **kwargs): def query(self, query, db_id, **kwargs):
db = self.get_kb_by_id(db_id)
distance_threshold = kwargs.get("distance_threshold", self.default_distance_threshold) distance_threshold = kwargs.get("distance_threshold", self.default_distance_threshold)
rerank_threshold = kwargs.get("rerank_threshold", self.default_rerank_threshold) rerank_threshold = kwargs.get("rerank_threshold", self.default_rerank_threshold)
@ -361,11 +359,19 @@ class KnowledgeBase:
def get_retrievers(self): def get_retrievers(self):
retrievers = {} retrievers = {}
for db in self.db_manager.get_all_databases(): for db in self.db_manager.get_all_databases():
retrievers[db["db_id"]] = { if self.check_embed_model(db["db_id"]):
"name": db["name"], retrievers[db["db_id"]] = {
"description": db["description"], "name": db["name"],
"retriever": self.get_retriever_by_db_id(db["db_id"]), "description": db["description"],
} "retriever": self.get_retriever_by_db_id(db["db_id"]),
"embed_model": db["embed_model"],
}
else:
logger.warning((
f"无法将知识库 {db['name']} 转换为 Tools, 因为向量模型不匹配,"
f"当前向量模型: {self.embed_model.embed_model_fullname}"
f"知识库向量模型: {db['embed_model']}"
))
return retrievers return retrievers
################################ ################################
@ -474,4 +480,9 @@ class KnowledgeBase:
def search_by_id(self, collection_name, id, output_fields=["id", "text"]): def search_by_id(self, collection_name, id, output_fields=["id", "text"]):
res = self.client.get(collection_name, id, output_fields=output_fields) res = self.client.get(collection_name, id, output_fields=output_fields)
return res return res
def check_embed_model(self, db_id):
db = self.db_manager.get_database_by_id(db_id)
return db["embed_model"] == self.embed_model.embed_model_fullname

48
src/core/operators.py Normal file
View File

@ -0,0 +1,48 @@
"""
这里面存放是 RAG 相关的一些组件
"""
from src.utils import prompts
class BaseOperator:
"""
基类
"""
template = None
def __init__(self):
pass
def call(self, **kwargs):
pass
def __call__(self, **kwargs):
"""
所有 RAG 相关组件的调用接口
"""
return self.call(**kwargs)
class HyDEOperator(BaseOperator):
"""
HyDE 重写查询
"""
template = prompts.HYDE_PROMPT_TEMPLATE
def __init__(self):
super().__init__()
@classmethod
def call(cls, model_callable, query, context_str, **kwargs):
"""
重写查询
Args:
model_callable: 模型调用函数
query: 查询
context_str: 上下文
"""
prompt = cls.template.format(query=query, context_str=context_str)
response = model_callable(prompt)
return response

View File

@ -2,6 +2,8 @@ from src import config, knowledge_base, graph_base
from src.models.rerank_model import get_reranker from src.models.rerank_model import get_reranker
from src.utils.logging_config import logger from src.utils.logging_config import logger
from src.models import select_model from src.models import select_model
from src.core.operators import HyDEOperator
class Retriever: class Retriever:
@ -151,8 +153,8 @@ class Retriever:
rewritten_query = model.predict(rewritten_query_prompt).content rewritten_query = model.predict(rewritten_query_prompt).content
if rewrite_query_span == "hyde": if rewrite_query_span == "hyde":
hy_doc = model.predict(rewritten_query).content res = HyDEOperator.call(model_callable=model.predict, query=query, context_str=history_query)
rewritten_query = f"{rewritten_query} {hy_doc}" rewritten_query = res.content
return rewritten_query return rewritten_query

View File

@ -56,4 +56,16 @@ keywords_prompt_template = """
返回的实体使用<->隔开关键词1<->关键词<->关键词3 返回的实体使用<->隔开关键词1<->关键词<->关键词3
不要改变关键词的语言 不要改变关键词的语言
<文本>{text}</文本> <文本>{text}</文本>
""" """
HYDE_PROMPT_TEMPLATE = (
"Please write a passage to answer the question\n"
"Try to include as many key details as possible.\n"
"\n"
"\n"
"{context_str}\n"
"\n"
"{query}\n"
"\n"
'Passage:\n'
)

View File

@ -165,7 +165,7 @@
<!-- 添加工具选择部分 --> <!-- 添加工具选择部分 -->
<a-form-item label="可用工具" name="tools" class="config-item"> <a-form-item label="可用工具" name="tools" class="config-item">
<p class="description">选择要启用的工具</p> <p class="description">选择要启用的工具retrieve 工具仅展现了与当前向量模型匹配的知识库详情请查看 docker 日志</p>
<a-form-item-rest> <a-form-item-rest>
<div class="tools-switches"> <div class="tools-switches">
<div v-for="tool in availableTools" :key="tool" class="tool-switch-item"> <div v-for="tool in availableTools" :key="tool" class="tool-switch-item">

View File

@ -5,18 +5,19 @@
description="知识型数据库,主要是非结构化的文本组成,使用向量检索使用。如果出现问题,可以检查 saves/data/database.json 查看配置。" description="知识型数据库,主要是非结构化的文本组成,使用向量检索使用。如果出现问题,可以检查 saves/data/database.json 查看配置。"
> >
<template #actions> <template #actions>
<a-button type="primary" @click="newDatabase.open=true">新建数据</a-button> <a-button type="primary" @click="newDatabase.open=true">新建知识</a-button>
</template> </template>
</HeaderComponent> </HeaderComponent>
<a-modal :open="newDatabase.open" title="新建数据库" @ok="createDatabase" @cancel="newDatabase.open=false"> <a-modal :open="newDatabase.open" title="新建知识库" @ok="createDatabase" @cancel="newDatabase.open=false">
<h3>数据库名称<span style="color: var(--error-color)">*</span></h3> <h3>知识库名称<span style="color: var(--error-color)">*</span></h3>
<a-input v-model:value="newDatabase.name" placeholder="新建数据库名称" /> <a-input v-model:value="newDatabase.name" placeholder="新建知识库名称" />
<h3 style="margin-top: 20px;">数据库描述</h3> <h3 style="margin-top: 20px;">知识库描述</h3>
<p style="color: var(--gray-700); font-size: 14px;">在智能体流程中这里的描述会作为工具的描述智能体会根据知识库的标题和描述来选择合适的工具所以这里描述的越详细智能体越容易选择到合适的工具</p>
<a-textarea <a-textarea
v-model:value="newDatabase.description" v-model:value="newDatabase.description"
placeholder="新建数据库描述" placeholder="新建知识库描述"
:auto-size="{ minRows: 2, maxRows: 5 }" :auto-size="{ minRows: 5, maxRows: 10 }"
/> />
<!-- <h3 style="margin-top: 20px;">向量维度</h3> <!-- <h3 style="margin-top: 20px;">向量维度</h3>
<p>必须与向量模型 {{ configStore.config.embed_model }} 一致</p> <p>必须与向量模型 {{ configStore.config.embed_model }} 一致</p>
@ -31,7 +32,7 @@
<div class="top"> <div class="top">
<div class="icon"><PlusOutlined /></div> <div class="icon"><PlusOutlined /></div>
<div class="info"> <div class="info">
<h3>新建数据</h3> <h3>新建知识</h3>
</div> </div>
</div> </div>
<p>导入您自己的文本数据或通过Webhook实时写入数据以增强 LLM 的上下文</p> <p>导入您自己的文本数据或通过Webhook实时写入数据以增强 LLM 的上下文</p>