ForcePilot/src/core/retriever.py

236 lines
8.7 KiB
Python
Raw Normal View History

from src.models.embedding import Reranker
from src.utils.logging_config import setup_logger
2024-09-06 21:30:38 +08:00
2024-07-20 20:30:32 +08:00
logger = setup_logger("server-common")
2024-07-16 23:12:35 +08:00
2024-07-10 12:56:53 +08:00
class Retriever:
def __init__(self, config, dbm, model):
2024-07-10 12:56:53 +08:00
self.config = config
self.dbm = dbm
self.model = model
2024-07-10 12:56:53 +08:00
2024-07-29 01:00:02 +08:00
if self.config.enable_reranker:
self.reranker = Reranker(config)
2024-07-17 18:52:20 +08:00
def retrieval(self, query, history, meta):
2024-07-10 12:56:53 +08:00
2024-08-08 19:51:13 +08:00
refs = {"query": query, "history": history, "meta": meta}
2024-09-14 02:45:13 +08:00
refs["model_name"] = self.config.model_name
2024-07-29 01:00:02 +08:00
refs["entities"] = self.reco_entities(query, history, refs)
refs["knowledge_base"] = self.query_knowledgebase(query, history, refs)
refs["graph_base"] = self.query_graph(query, history, refs)
2024-07-10 12:56:53 +08:00
2024-07-17 18:52:20 +08:00
return refs
2024-07-10 12:56:53 +08:00
2024-07-17 18:52:20 +08:00
def construct_query(self, query, refs, meta):
2024-09-18 21:03:12 +08:00
if not refs or len(refs) == 0:
2024-07-10 12:56:53 +08:00
return query
2024-09-18 21:03:12 +08:00
external_parts = []
2024-07-10 12:56:53 +08:00
2024-07-29 01:00:02 +08:00
# 解析知识库的结果
2024-09-18 21:03:12 +08:00
kb_res = refs.get("knowledge_base", {}).get("results", [])
if kb_res:
kb_text = "\n".join(f"{r['id']}: {r['entity']['text']}" for r in kb_res)
external_parts.extend(["知识库信息:", kb_text])
2024-07-10 12:56:53 +08:00
2024-07-29 01:00:02 +08:00
# 解析图数据库的结果
2024-09-18 21:03:12 +08:00
db_res = refs.get("graph_base", {}).get("results", {})
if db_res.get("nodes") and len(db_res["nodes"]) > 0:
2024-09-06 21:30:38 +08:00
db_text = "\n".join(
2024-09-18 21:03:12 +08:00
[f"{edge['source_name']}{edge['target_name']}的关系是{edge['type']}" for edge in db_res.get("edges", [])]
2024-09-06 21:30:38 +08:00
)
2024-09-18 21:03:12 +08:00
external_parts.extend(["图数据库信息:", db_text])
2024-07-17 18:52:20 +08:00
2024-07-29 01:00:02 +08:00
# 构造查询
2024-09-25 13:46:23 +08:00
from src.utils.prompts import knowbase_qa_template
2024-09-18 21:03:12 +08:00
if external_parts and len(external_parts) > 0:
external = "\n\n".join(external_parts)
2024-09-25 13:46:23 +08:00
query = knowbase_qa_template.format(external=external, query=query)
2024-07-10 12:56:53 +08:00
return query
def query_classification(self, query):
"""判断是否需要查询
- 对于完全基于用户给定信息的任务称之为足够sufficient不需要检索
- 否则称之为不足insufficient可能需要检索
"""
raise NotImplementedError
2024-07-29 01:00:02 +08:00
def query_graph(self, query, history, refs):
2024-07-16 23:12:35 +08:00
# res = model.predict("qiansdgsa, dasdh ashdsakjdk ak ").content
2024-07-17 18:50:01 +08:00
results = []
2024-07-31 20:22:05 +08:00
if refs["meta"].get("use_graph") and self.config.enable_knowledge_base:
2024-07-29 01:00:02 +08:00
for entity in refs["entities"]:
result = self.dbm.graph_base.query_by_vector(entity)
2024-07-20 20:30:32 +08:00
if result != []:
results.extend(result)
2024-07-18 02:46:58 +08:00
return {"results": self.format_query_results(results)}
2024-07-16 23:12:35 +08:00
2024-07-29 01:00:02 +08:00
def query_knowledgebase(self, query, history, refs):
"""查询知识库"""
2024-07-17 18:52:20 +08:00
2024-07-17 19:10:41 +08:00
kb_res = []
2024-07-31 20:22:05 +08:00
final_res = []
2024-09-18 21:03:12 +08:00
db_name = refs["meta"].get("db_name")
if not db_name or not self.config.enable_knowledge_base:
2024-09-06 21:30:38 +08:00
return {
"results": final_res,
"all_results": kb_res,
"rw_query": query,
"message": "Knowledge base is disabled",
}
2024-08-08 19:51:13 +08:00
2024-08-25 20:29:24 +08:00
rw_query = self.rewrite_query(query, history, refs)
2024-08-08 19:51:13 +08:00
2024-08-25 20:29:24 +08:00
kb = self.dbm.metaname2db[db_name]
2024-09-13 15:25:01 +08:00
logger.debug(f"{refs['meta']=}")
2024-07-17 18:52:20 +08:00
2024-09-18 21:03:12 +08:00
meta = refs["meta"]
max_query_count = meta.get("maxQueryCount", 10)
rerank_threshold = meta.get("rerankThreshold", 0.1)
distance_threshold = meta.get("distanceThreshold", 0)
top_k = meta.get("topK", 5)
2024-09-09 17:07:03 +08:00
all_kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count)
for r in all_kb_res:
2024-08-25 20:29:24 +08:00
r["file"] = kb.id2file(r["entity"]["file_id"])
2024-09-09 17:07:03 +08:00
# use distance threshold to filter results
2024-09-18 21:03:12 +08:00
if meta.get("mode") == "search":
2024-09-14 02:45:13 +08:00
kb_res = all_kb_res
else:
kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold]
2024-09-09 17:07:03 +08:00
2024-08-25 20:29:24 +08:00
if self.config.enable_reranker:
for r in kb_res:
2024-09-13 15:25:01 +08:00
r["rerank_score"] = self.reranker.compute_score([rw_query, r["entity"]["text"]], normalize=True)[0]
2024-08-25 20:29:24 +08:00
kb_res.sort(key=lambda x: x["rerank_score"], reverse=True)
2024-09-09 17:07:03 +08:00
kb_res = [_res for _res in kb_res if _res["rerank_score"] > rerank_threshold]
2024-08-08 19:51:13 +08:00
2024-09-09 17:07:03 +08:00
kb_res = kb_res[:top_k]
2024-09-09 17:07:03 +08:00
return {"results": kb_res, "all_results": all_kb_res, "rw_query": rw_query}
2024-07-17 18:52:20 +08:00
2024-07-29 01:00:02 +08:00
def rewrite_query(self, query, history, refs):
2024-07-10 12:56:53 +08:00
"""重写查询"""
2024-11-23 20:53:51 +08:00
if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置
rewrite_query_span = refs["meta"].get("use_rewrite_query", "off")
else:
rewrite_query_span = refs["meta"]["config"].get("use_rewrite_query", "off")
2024-09-09 17:07:03 +08:00
if rewrite_query_span == "off":
2024-07-17 18:50:01 +08:00
rewritten_query = query
else:
2024-07-29 01:00:02 +08:00
from src.utils.prompts import rewritten_query_prompt_template
2024-09-06 21:30:38 +08:00
history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else ""
2024-08-08 19:51:13 +08:00
rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query)
rewritten_query = self.model.predict(rewritten_query_prompt).content
2024-07-17 18:50:01 +08:00
2024-09-09 17:07:03 +08:00
if rewrite_query_span == "hyde":
2024-08-08 19:51:13 +08:00
hy_doc = self.model.predict(rewritten_query).content
rewritten_query = f"{rewritten_query} {hy_doc}"
2024-07-29 01:00:02 +08:00
return rewritten_query
2024-07-21 18:15:28 +08:00
2024-07-29 01:00:02 +08:00
def reco_entities(self, query, history, refs):
"""识别句子中的实体"""
query = refs.get("rewritten_query", query)
entities = []
if refs["meta"].get("use_graph"):
2024-09-18 21:03:12 +08:00
from src.utils.prompts import entity_extraction_prompt_template as entity_template
from src.utils.prompts import keywords_prompt_template as entity_template
2024-09-06 21:30:38 +08:00
2024-09-18 21:03:12 +08:00
entity_extraction_prompt = entity_template.format(text=query)
entities = self.model.predict(entity_extraction_prompt).content.split("<->")
# entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)]
2024-07-29 01:00:02 +08:00
return entities
2024-07-10 12:56:53 +08:00
2024-09-06 21:30:38 +08:00
def _extract_relationship_info(self, relationship, source_name, target_name):
"""
提取关系信息并返回格式化的节点和边信息
"""
rel_id = relationship.element_id
nodes = relationship.nodes
if len(nodes) != 2:
return None, None
source, target = nodes
source_id = source.element_id
target_id = target.element_id
relationship_type = relationship._properties.get("type", "unknown")
if relationship_type == "unknown":
relationship_type = relationship.type
edge_info = {
"id": rel_id,
"type": relationship_type,
"source_id": source_id,
"target_id": target_id,
"source_name": source_name,
"target_name": target_name,
}
node_info = [
{"id": source_id, "name": source_name},
{"id": target_id, "name": target_name},
]
return node_info, edge_info
def format_general_results(self, results):
2024-09-06 12:54:17 +08:00
formatted_results = {"nodes": [], "edges": []}
for item in results:
relationship = item[1]
2024-09-06 21:30:38 +08:00
source_name = item[0]._properties.get("name", "unknown")
target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "unknown"
2024-09-06 12:54:17 +08:00
2024-09-06 21:30:38 +08:00
node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name)
if node_info is None or edge_info is None:
continue
2024-09-06 12:54:17 +08:00
2024-09-06 21:30:38 +08:00
for node in node_info:
if node["id"] not in [n["id"] for n in formatted_results["nodes"]]:
formatted_results["nodes"].append(node)
2024-09-06 12:54:17 +08:00
2024-09-06 21:30:38 +08:00
formatted_results["edges"].append(edge_info)
2024-09-06 12:54:17 +08:00
return formatted_results
2024-07-29 01:00:02 +08:00
def format_query_results(self, results):
2024-07-18 02:46:58 +08:00
formatted_results = {"nodes": [], "edges": []}
2024-07-20 20:30:32 +08:00
node_dict = {}
2024-07-20 20:30:32 +08:00
for item in results:
2024-09-18 21:03:12 +08:00
if not isinstance(item[1], list) or not item[1]:
2024-07-29 01:00:02 +08:00
continue
relationship = item[1][0]
2024-09-06 21:30:38 +08:00
source_name = item[0]
target_name = item[2] if len(item) > 2 else "unknown"
node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name)
if node_info is None or edge_info is None:
2024-07-29 01:00:02 +08:00
continue
2024-09-18 21:03:12 +08:00
node_dict.update({node["id"]: node for node in node_info})
2024-09-06 21:30:38 +08:00
formatted_results["edges"].append(edge_info)
2024-07-20 20:30:32 +08:00
formatted_results["nodes"] = list(node_dict.values())
2024-07-18 02:46:58 +08:00
return formatted_results
2024-07-17 18:52:20 +08:00
def __call__(self, query, history, meta):
refs = self.retrieval(query, history, meta)
query = self.construct_query(query, refs, meta)
2024-09-06 21:30:38 +08:00
return query, refs