优化知识库查询以及工具定义方法
This commit is contained in:
parent
f11f719b4e
commit
380e7e38d6
@ -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="查询的关键词,查询的时候,应该尽量以关键词的形式进行查询,不要使用复杂的句子。\n.")
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
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
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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">
|
||||||
|
|||||||
@ -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>
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user