修复知识图谱检索bug

This commit is contained in:
Wenjie Zhang 2025-04-13 21:06:23 +08:00
parent 1d20db6d15
commit 8323279ff4
2 changed files with 68 additions and 34 deletions

View File

@ -206,19 +206,12 @@ class GraphDatabase:
""" """
tx.run(query) tx.run(query)
def query_node(self, entity_name, hops=2, **kwargs): def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, **kwargs):
# TODO 添加判断节点数量为 0 停止检索 # TODO 添加判断节点数量为 0 停止检索
# 判断是否启动 # 判断是否启动
if not self.is_running(): if not self.is_running():
raise Exception("图数据库未启动") raise Exception("图数据库未启动")
logger.debug(f"Query graph node {entity_name} with {hops=}")
if kwargs.get("exact_match"):
raise NotImplemented("not implement for `exact_match`")
else:
return self.query_by_vector(entity_name=entity_name, **kwargs)
def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=5):
self.use_database(kgdb_name) self.use_database(kgdb_name)
def _index_exists(tx, index_name): def _index_exists(tx, index_name):
"""检查索引是否存在""" """检查索引是否存在"""
@ -251,7 +244,7 @@ class GraphDatabase:
raise e raise e
# 筛选出分数高于阈值的实体 # 筛选出分数高于阈值的实体
qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold] qualified_entities = [result[0] for result in results[:max_entities] if result[1] > threshold]
logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}") logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}")
# 对每个合格的实体进行查询 # 对每个合格的实体进行查询
@ -262,70 +255,102 @@ class GraphDatabase:
return all_query_results return all_query_results
def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2): def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2, limit=100):
"""查询指定实体三元组信息(无向关系)""" """查询指定实体三元组信息(无向关系)"""
self.use_database(kgdb_name) if not entity_name:
def query(tx, entity_name, hops): logger.warning("实体名称为空")
result = tx.run(f""" return []
MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m)
RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m
""", entity_name=entity_name)
return result.values()
with self.driver.session() as session: self.use_database(kgdb_name)
return session.execute_read(query, entity_name, hops)
def query(tx, entity_name, hops, limit):
try:
query_str = f"""
MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m)
RETURN n AS n, r, m AS m
LIMIT $limit
"""
result = tx.run(query_str, entity_name=entity_name, limit=limit)
if not result:
logger.info(f"未找到实体 {entity_name} 的相关信息")
return []
values = result.values()
# 安全地处理embedding属性
values = clean_triples_embedding(values)
return values
except Exception as e:
logger.error(f"查询实体 {entity_name} 失败: {str(e)}")
return []
try:
with self.driver.session() as session:
return session.execute_read(query, entity_name, hops, limit)
except Exception as e:
logger.error(f"数据库会话异常: {str(e)}")
return []
def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2): def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2):
"""查询图数据库中所有三元组信息""" """查询图数据库中所有三元组信息 NEVER USE"""
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, hops): def query(tx, hops):
result = tx.run(f""" result = tx.run(f"""
MATCH (n)-[r*1..{hops}]->(m) MATCH (n)-[r*1..{hops}]->(m)
RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m RETURN n AS n, r, m AS m
""") """)
return result.values() values = result.values()
values = clean_triples_embedding(values)
return values
with self.driver.session() as session: with self.driver.session() as session:
return session.execute_read(query, hops) return session.execute_read(query, hops)
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2): def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
"""查询指定关系三元组信息""" """查询指定关系三元组信息 NEVER USE"""
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, relationship_type, hops): def query(tx, relationship_type, hops):
result = tx.run(f""" result = tx.run(f"""
MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m) MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m)
RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m RETURN n AS n, r, m AS m
""") """)
return result.values() values = result.values()
values = clean_triples_embedding(values)
return values
with self.driver.session() as session: with self.driver.session() as session:
return session.execute_read(query, relationship_type, hops) return session.execute_read(query, relationship_type, hops)
def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2): def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2):
"""模糊查询""" """模糊查询 NEVER USE"""
self.use_database(kgdb_name) self.use_database(kgdb_name)
def query(tx, keyword, hops): def query(tx, keyword, hops):
result = tx.run(f""" result = tx.run(f"""
MATCH (n:Entity) MATCH (n:Entity)
WHERE n.name CONTAINS $keyword WHERE n.name CONTAINS $keyword
MATCH (n)-[r*1..{hops}]->(m) MATCH (n)-[r*1..{hops}]->(m)
RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m RETURN n AS n, r, m AS m
""", keyword=keyword) """, keyword=keyword)
return result.values() values = result.values()
values = clean_triples_embedding(values)
return values
with self.driver.session() as session: with self.driver.session() as session:
return session.execute_read(query, keyword, hops) return session.execute_read(query, keyword, hops)
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
"""查询指定节点的详细信息返回信息""" """查询指定节点的详细信息返回信息 NEVER USE"""
self.use_database(kgdb_name) # 切换到指定数据库 self.use_database(kgdb_name) # 切换到指定数据库
def query(tx, node_name, hops): def query(tx, node_name, hops):
result = tx.run(f""" result = tx.run(f"""
MATCH (n {{name: $node_name}}) MATCH (n {{name: $node_name}})
OPTIONAL MATCH (n)-[r*1..{hops}]->(m) OPTIONAL MATCH (n)-[r*1..{hops}]->(m)
RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m RETURN n AS n, r, m AS m
""", node_name=node_name) """, node_name=node_name)
return result.values() values = result.values()
values = clean_triples_embedding(values)
return values
with self.driver.session() as session: with self.driver.session() as session:
return session.execute_read(query, node_name, hops) return session.execute_read(query, node_name, hops)
@ -473,6 +498,14 @@ class GraphDatabase:
return count return count
def clean_triples_embedding(triples):
for item in triples:
if hasattr(item[0], '_properties'):
item[0]._properties['embedding'] = None
if hasattr(item[2], '_properties'):
item[2]._properties['embedding'] = None
return triples
if __name__ == "__main__": if __name__ == "__main__":
pass pass

View File

@ -1,10 +1,11 @@
import traceback
from src import config, knowledge_base, graph_base 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 from src.core.operators import HyDEOperator
class Retriever: class Retriever:
def __init__(self): def __init__(self):
@ -80,7 +81,7 @@ class Retriever:
for entity in refs["entities"]: for entity in refs["entities"]:
if entity == "": if entity == "":
continue continue
result = graph_base.query_by_vector(entity) result = graph_base.query_node(entity)
if result != []: if result != []:
results.extend(result) results.extend(result)
return {"results": self.format_query_results(results)} return {"results": self.format_query_results(results)}
@ -256,7 +257,7 @@ class Retriever:
# 添加边 # 添加边
formatted_results["edges"].append(edge_info) formatted_results["edges"].append(edge_info)
except Exception as e: except Exception as e:
logger.error(f"处理关系时出错: {e}, 关系: {relationship}") logger.error(f"处理关系时出错: {e}, 关系: {relationship}, {traceback.format_exc()}")
continue continue
# 将节点字典转换为列表 # 将节点字典转换为列表