fix graph bugs
This commit is contained in:
parent
0d455bcc3d
commit
b7e65fe108
@ -1,5 +1,6 @@
|
|||||||
from src.models.embedding import Reranker
|
from src.models.embedding import Reranker
|
||||||
from src.utils.logging_config import setup_logger
|
from src.utils.logging_config import setup_logger
|
||||||
|
|
||||||
logger = setup_logger("server-common")
|
logger = setup_logger("server-common")
|
||||||
|
|
||||||
|
|
||||||
@ -37,12 +38,14 @@ class Retriever:
|
|||||||
# 解析图数据库的结果
|
# 解析图数据库的结果
|
||||||
db_res = refs.get("graph_base").get("results", [])
|
db_res = refs.get("graph_base").get("results", [])
|
||||||
if len(db_res["nodes"]) > 0:
|
if len(db_res["nodes"]) > 0:
|
||||||
db_text = '\n'.join([f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res['edges']])
|
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}"
|
external += f"图数据库信息: \n\n{db_text}"
|
||||||
|
|
||||||
# 构造查询
|
# 构造查询
|
||||||
if len(external) > 0:
|
if len(external) > 0:
|
||||||
query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}"
|
query = f"参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答问题。\n\n问题:{query}\n\n回答:"
|
||||||
|
|
||||||
return query
|
return query
|
||||||
|
|
||||||
@ -70,7 +73,12 @@ class Retriever:
|
|||||||
kb_res = []
|
kb_res = []
|
||||||
final_res = []
|
final_res = []
|
||||||
if not refs["meta"].get("db_name") or not self.config.enable_knowledge_base:
|
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"}
|
return {
|
||||||
|
"results": final_res,
|
||||||
|
"all_results": kb_res,
|
||||||
|
"rw_query": query,
|
||||||
|
"message": "Knowledge base is disabled",
|
||||||
|
}
|
||||||
|
|
||||||
rw_query = self.rewrite_query(query, history, refs)
|
rw_query = self.rewrite_query(query, history, refs)
|
||||||
|
|
||||||
@ -85,7 +93,7 @@ class Retriever:
|
|||||||
if self.config.enable_reranker:
|
if self.config.enable_reranker:
|
||||||
RERANK_THRESHOLD = 0.001
|
RERANK_THRESHOLD = 0.001
|
||||||
for r in kb_res:
|
for r in kb_res:
|
||||||
r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True)
|
r["rerank_score"] = self.reranker.compute_score([rw_query, r["entity"]["text"]], normalize=True)
|
||||||
kb_res.sort(key=lambda x: x["rerank_score"], reverse=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]
|
final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD]
|
||||||
|
|
||||||
@ -101,7 +109,8 @@ class Retriever:
|
|||||||
rewritten_query = query
|
rewritten_query = query
|
||||||
else:
|
else:
|
||||||
from src.utils.prompts import rewritten_query_prompt_template
|
from src.utils.prompts import rewritten_query_prompt_template
|
||||||
history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else ""
|
|
||||||
|
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)
|
rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query)
|
||||||
rewritten_query = self.model.predict(rewritten_query_prompt).content
|
rewritten_query = self.model.predict(rewritten_query_prompt).content
|
||||||
|
|
||||||
@ -118,54 +127,70 @@ class Retriever:
|
|||||||
entities = []
|
entities = []
|
||||||
if refs["meta"].get("use_graph"):
|
if refs["meta"].get("use_graph"):
|
||||||
from src.utils.prompts import entity_extraction_prompt_template
|
from src.utils.prompts import entity_extraction_prompt_template
|
||||||
|
|
||||||
entity_extraction_prompt = entity_extraction_prompt_template.format(text=query)
|
entity_extraction_prompt = entity_extraction_prompt_template.format(text=query)
|
||||||
entities = self.model.predict(entity_extraction_prompt).content.split(",")
|
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)]
|
entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)]
|
||||||
|
|
||||||
return entities
|
return entities
|
||||||
|
|
||||||
def foramt_general_results(self, results):
|
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):
|
||||||
logger.debug(f"Formatting general results: {results}")
|
logger.debug(f"Formatting general results: {results}")
|
||||||
formatted_results = {"nodes": [], "edges": []}
|
formatted_results = {"nodes": [], "edges": []}
|
||||||
|
|
||||||
for item in results:
|
for item in results:
|
||||||
relationship = item[1]
|
relationship = item[1]
|
||||||
rel_id = relationship.element_id
|
source_name = item[0]._properties.get("name", "unknown")
|
||||||
nodes = relationship.nodes
|
target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "unknown"
|
||||||
if len(nodes) != 2:
|
|
||||||
|
node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name)
|
||||||
|
if node_info is None or edge_info is None:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
source, target = nodes
|
for node in node_info:
|
||||||
|
if node["id"] not in [n["id"] for n in formatted_results["nodes"]]:
|
||||||
|
formatted_results["nodes"].append(node)
|
||||||
|
|
||||||
source_id = source.element_id
|
formatted_results["edges"].append(edge_info)
|
||||||
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
|
return formatted_results
|
||||||
|
|
||||||
def format_query_results(self, results):
|
def format_query_results(self, results):
|
||||||
logger.debug(f"Formatting query results: {results}")
|
logger.debug(f"Formatting query results: {results}")
|
||||||
formatted_results = {"nodes": [], "edges": []}
|
formatted_results = {"nodes": [], "edges": []}
|
||||||
|
|
||||||
node_dict = {}
|
node_dict = {}
|
||||||
|
|
||||||
for item in results:
|
for item in results:
|
||||||
@ -173,35 +198,18 @@ class Retriever:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
relationship = item[1][0]
|
relationship = item[1][0]
|
||||||
rel_id = relationship.element_id
|
source_name = item[0]
|
||||||
nodes = relationship.nodes
|
target_name = item[2] if len(item) > 2 else "unknown"
|
||||||
if len(nodes) != 2:
|
|
||||||
|
node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name)
|
||||||
|
if node_info is None or edge_info is None:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
source, target = nodes
|
for node in node_info:
|
||||||
|
if node["id"] not in node_dict:
|
||||||
|
node_dict[node["id"]] = node
|
||||||
|
|
||||||
source_id = source.element_id
|
formatted_results["edges"].append(edge_info)
|
||||||
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
|
|
||||||
})
|
|
||||||
|
|
||||||
formatted_results["nodes"] = list(node_dict.values())
|
formatted_results["nodes"] = list(node_dict.values())
|
||||||
|
|
||||||
@ -210,4 +218,4 @@ class Retriever:
|
|||||||
def __call__(self, query, history, meta):
|
def __call__(self, query, history, meta):
|
||||||
refs = self.retrieval(query, history, meta)
|
refs = self.retrieval(query, history, meta)
|
||||||
query = self.construct_query(query, refs, meta)
|
query = self.construct_query(query, refs, meta)
|
||||||
return query, refs
|
return query, refs
|
||||||
|
|||||||
@ -135,7 +135,7 @@ def get_graph_nodes():
|
|||||||
|
|
||||||
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
||||||
result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
||||||
return jsonify({'result': startup.retriever.foramt_general_results(result), 'message': 'success'}), 200
|
return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200
|
||||||
|
|
||||||
@db.route('/graph/add', methods=['POST'])
|
@db.route('/graph/add', methods=['POST'])
|
||||||
def add_graph_entity():
|
def add_graph_entity():
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user