2024-07-28 16:16:52 +08:00
|
|
|
|
from src.models.embedding import Reranker
|
|
|
|
|
|
from src.utils.logging_config import setup_logger
|
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:
|
|
|
|
|
|
|
2024-07-22 00:00:54 +08:00
|
|
|
|
def __init__(self, config, dbm, model):
|
2024-07-10 12:56:53 +08:00
|
|
|
|
self.config = config
|
2024-07-22 00:00:54 +08:00
|
|
|
|
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-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-07-10 12:56:53 +08:00
|
|
|
|
if len(refs) == 0:
|
|
|
|
|
|
return query
|
|
|
|
|
|
|
|
|
|
|
|
external = ""
|
|
|
|
|
|
|
2024-07-29 01:00:02 +08:00
|
|
|
|
# 解析知识库的结果
|
2024-07-17 18:52:20 +08:00
|
|
|
|
kb_res = refs.get("knowledge_base").get("results", [])
|
|
|
|
|
|
if len(kb_res) > 0:
|
2024-07-10 12:56:53 +08:00
|
|
|
|
kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res])
|
|
|
|
|
|
external += f"知识库信息: \n\n{kb_text}"
|
|
|
|
|
|
|
2024-07-29 01:00:02 +08:00
|
|
|
|
# 解析图数据库的结果
|
2024-07-20 20:30:32 +08:00
|
|
|
|
db_res = refs.get("graph_base").get("results", [])
|
2024-07-25 20:30:28 +08:00
|
|
|
|
if len(db_res["nodes"]) > 0:
|
2024-07-20 20:30:32 +08:00
|
|
|
|
db_text = '\n'.join([f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res['edges']])
|
|
|
|
|
|
external += f"图数据库信息: \n\n{db_text}"
|
2024-07-17 18:52:20 +08:00
|
|
|
|
|
2024-07-29 01:00:02 +08:00
|
|
|
|
# 构造查询
|
2024-07-10 12:56:53 +08:00
|
|
|
|
if len(external) > 0:
|
|
|
|
|
|
query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}"
|
|
|
|
|
|
|
|
|
|
|
|
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 != []:
|
2024-07-21 18:17:25 +08:00
|
|
|
|
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-08-25 20:29:24 +08:00
|
|
|
|
if not refs["meta"].get("db_name") or not self.config.enable_knowledge_base:
|
|
|
|
|
|
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
|
|
|
|
db_name = refs["meta"]["db_name"]
|
|
|
|
|
|
kb = self.dbm.metaname2db[db_name]
|
|
|
|
|
|
limit = refs["meta"].get("queryCount", 10)
|
2024-07-17 18:52:20 +08:00
|
|
|
|
|
2024-08-25 20:29:24 +08:00
|
|
|
|
kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit)
|
|
|
|
|
|
for r in kb_res:
|
|
|
|
|
|
r["file"] = kb.id2file(r["entity"]["file_id"])
|
|
|
|
|
|
|
|
|
|
|
|
if self.config.enable_reranker:
|
2024-09-06 12:54:17 +08:00
|
|
|
|
RERANK_THRESHOLD = 0.001
|
2024-08-25 20:29:24 +08:00
|
|
|
|
for r in kb_res:
|
|
|
|
|
|
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
|
|
|
|
|
|
kb_res.sort(key=lambda x: x["rerank_score"], reverse=True)
|
|
|
|
|
|
final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD]
|
2024-08-08 19:51:13 +08:00
|
|
|
|
|
2024-08-25 20:29:24 +08:00
|
|
|
|
else:
|
|
|
|
|
|
final_res = kb_res[:5]
|
2024-07-28 16:16:52 +08:00
|
|
|
|
|
2024-08-08 19:51:13 +08:00
|
|
|
|
return {"results": final_res, "all_results": 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-08-08 19:51:13 +08:00
|
|
|
|
rewrite_query_span = refs["meta"].get("rewrite_query", None)
|
|
|
|
|
|
if rewrite_query_span is None or 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-08-08 19:51:13 +08:00
|
|
|
|
history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else ""
|
|
|
|
|
|
rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
rewritten_query = self.model.predict(rewritten_query_prompt).content
|
2024-07-17 18:50:01 +08:00
|
|
|
|
|
2024-08-08 19:51:13 +08:00
|
|
|
|
if rewrite_query_span == "HyDE":
|
|
|
|
|
|
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"):
|
|
|
|
|
|
from src.utils.prompts import entity_extraction_prompt_template
|
|
|
|
|
|
entity_extraction_prompt = entity_extraction_prompt_template.format(text=query)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
entities = self.model.predict(entity_extraction_prompt).content.split(",")
|
2024-07-21 18:15:28 +08:00
|
|
|
|
entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)]
|
2024-07-17 18:55:04 +08:00
|
|
|
|
|
2024-07-29 01:00:02 +08:00
|
|
|
|
return entities
|
2024-07-10 12:56:53 +08:00
|
|
|
|
|
2024-09-06 12:54:17 +08:00
|
|
|
|
def foramt_general_results(self, results):
|
|
|
|
|
|
logger.debug(f"Formatting general results: {results}")
|
|
|
|
|
|
formatted_results = {"nodes": [], "edges": []}
|
|
|
|
|
|
|
|
|
|
|
|
for item in results:
|
|
|
|
|
|
relationship = item[1]
|
|
|
|
|
|
rel_id = relationship.element_id
|
|
|
|
|
|
nodes = relationship.nodes
|
|
|
|
|
|
if len(nodes) != 2:
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
source, target = nodes
|
|
|
|
|
|
|
|
|
|
|
|
source_id = source.element_id
|
|
|
|
|
|
target_id = target.element_id
|
|
|
|
|
|
source_name = source._properties.get('name', 'unknown')
|
|
|
|
|
|
target_name = target._properties.get('name', 'unknown')
|
|
|
|
|
|
|
|
|
|
|
|
if source_id not in formatted_results["nodes"]:
|
|
|
|
|
|
formatted_results["nodes"].append({"id": source_id, "name": source_name})
|
|
|
|
|
|
if target_id not in formatted_results["nodes"]:
|
|
|
|
|
|
formatted_results["nodes"].append({"id": target_id, "name": target_name})
|
|
|
|
|
|
|
|
|
|
|
|
relationship_type = relationship._properties.get('type', 'unknown')
|
|
|
|
|
|
if relationship_type == 'unknown':
|
|
|
|
|
|
relationship_type = relationship.type
|
|
|
|
|
|
|
|
|
|
|
|
formatted_results["edges"].append({
|
|
|
|
|
|
"id": rel_id,
|
|
|
|
|
|
"type": relationship_type,
|
|
|
|
|
|
"source_id": source_id,
|
|
|
|
|
|
"target_id": target_id,
|
|
|
|
|
|
"source_name": source_name,
|
|
|
|
|
|
"target_name": target_name
|
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
|
|
return formatted_results
|
|
|
|
|
|
|
2024-07-29 01:00:02 +08:00
|
|
|
|
def format_query_results(self, results):
|
2024-09-06 12:54:17 +08:00
|
|
|
|
logger.debug(f"Formatting query results: {results}")
|
2024-07-18 02:46:58 +08:00
|
|
|
|
formatted_results = {"nodes": [], "edges": []}
|
2024-07-21 18:17:25 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
node_dict = {}
|
2024-07-21 18:17:25 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
for item in results:
|
2024-07-29 01:00:02 +08:00
|
|
|
|
if not isinstance(item[1], list) or len(item[1]) == 0:
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
relationship = item[1][0]
|
|
|
|
|
|
rel_id = relationship.element_id
|
|
|
|
|
|
nodes = relationship.nodes
|
|
|
|
|
|
if len(nodes) != 2:
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
source, target = nodes
|
|
|
|
|
|
|
|
|
|
|
|
source_id = source.element_id
|
|
|
|
|
|
target_id = target.element_id
|
|
|
|
|
|
source_name = item[0]
|
|
|
|
|
|
target_name = item[2] if len(item) > 2 else 'unknown'
|
|
|
|
|
|
|
|
|
|
|
|
if source_id not in node_dict:
|
|
|
|
|
|
node_dict[source_id] = {"id": source_id, "name": source_name}
|
|
|
|
|
|
if target_id not in node_dict:
|
|
|
|
|
|
node_dict[target_id] = {"id": target_id, "name": target_name}
|
|
|
|
|
|
|
|
|
|
|
|
relationship_type = relationship._properties.get('type', 'unknown')
|
|
|
|
|
|
if relationship_type == 'unknown':
|
|
|
|
|
|
relationship_type = relationship.type
|
|
|
|
|
|
|
|
|
|
|
|
formatted_results["edges"].append({
|
|
|
|
|
|
"id": rel_id,
|
|
|
|
|
|
"type": relationship_type,
|
|
|
|
|
|
"source_id": source_id,
|
|
|
|
|
|
"target_id": target_id,
|
|
|
|
|
|
"source_name": source_name,
|
|
|
|
|
|
"target_name": target_name
|
|
|
|
|
|
})
|
2024-07-21 18:17:25 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
formatted_results["nodes"] = list(node_dict.values())
|
2024-07-21 18:17:25 +08:00
|
|
|
|
|
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-07-14 16:42:38 +08:00
|
|
|
|
return query, refs
|