commit
1d20db6d15
@ -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
|
||||||
|
|||||||
@ -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
48
src/core/operators.py
Normal 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
|
||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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'
|
||||||
|
)
|
||||||
|
|||||||
@ -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