2024-07-14 23:59:52 +08:00
|
|
|
|
import os
|
2024-07-15 18:58:34 +08:00
|
|
|
|
import json
|
2025-03-07 01:05:50 +08:00
|
|
|
|
import warnings
|
2025-03-19 01:21:51 +08:00
|
|
|
|
import traceback
|
2025-03-07 01:05:50 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
import torch
|
2024-07-16 18:14:27 +08:00
|
|
|
|
from neo4j import GraphDatabase as GD
|
2024-07-20 20:30:32 +08:00
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
from src import config
|
2025-02-28 00:38:59 +08:00
|
|
|
|
from src.utils import logger
|
2024-07-24 20:23:31 +08:00
|
|
|
|
|
2024-09-06 12:54:17 +08:00
|
|
|
|
warnings.filterwarnings("ignore", category=UserWarning)
|
2024-07-24 20:23:31 +08:00
|
|
|
|
|
2024-09-09 17:07:03 +08:00
|
|
|
|
|
2024-07-24 18:32:14 +08:00
|
|
|
|
UIE_MODEL = None
|
2024-07-14 23:59:52 +08:00
|
|
|
|
|
2024-07-16 18:14:27 +08:00
|
|
|
|
class GraphDatabase:
|
2025-03-20 19:51:46 +08:00
|
|
|
|
def __init__(self):
|
2024-07-16 18:14:27 +08:00
|
|
|
|
self.driver = None
|
|
|
|
|
|
self.files = []
|
|
|
|
|
|
self.status = "closed"
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.kgdb_name = "neo4j"
|
2025-02-23 16:39:52 +08:00
|
|
|
|
self.embed_model_name = None
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.work_dir = os.path.join(config.save_dir, "knowledge_graph", self.kgdb_name)
|
2025-02-28 00:38:59 +08:00
|
|
|
|
os.makedirs(self.work_dir, exist_ok=True)
|
|
|
|
|
|
|
2025-02-28 02:40:34 +08:00
|
|
|
|
# 尝试加载已保存的图数据库信息
|
2025-03-11 14:18:48 +08:00
|
|
|
|
if not self.load_graph_info():
|
2025-02-28 02:40:34 +08:00
|
|
|
|
logger.info(f"未找到已保存的图数据库信息,将创建新的配置")
|
|
|
|
|
|
|
2025-02-28 00:38:59 +08:00
|
|
|
|
self.start()
|
2024-07-16 18:14:27 +08:00
|
|
|
|
|
|
|
|
|
|
def start(self):
|
2025-03-20 19:51:46 +08:00
|
|
|
|
if not config.enable_knowledge_graph or not config.enable_knowledge_base:
|
|
|
|
|
|
return
|
2024-09-11 01:08:13 +08:00
|
|
|
|
uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687")
|
|
|
|
|
|
username = os.environ.get("NEO4J_USERNAME", "neo4j")
|
|
|
|
|
|
password = os.environ.get("NEO4J_PASSWORD", "0123456789")
|
|
|
|
|
|
logger.info(f"Connecting to Neo4j at {uri}/{self.kgdb_name}")
|
2024-10-08 22:16:17 +08:00
|
|
|
|
try:
|
|
|
|
|
|
self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password))
|
|
|
|
|
|
self.status = "open"
|
2025-03-20 19:51:46 +08:00
|
|
|
|
logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_graph_info(self.kgdb_name)}")
|
2025-02-28 02:40:34 +08:00
|
|
|
|
# 连接成功后保存图数据库信息
|
2025-03-20 19:51:46 +08:00
|
|
|
|
self.save_graph_info(self.kgdb_name)
|
2024-10-08 22:16:17 +08:00
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}")
|
|
|
|
|
|
self.config.enable_knowledge_graph = False
|
|
|
|
|
|
|
2024-07-15 18:58:34 +08:00
|
|
|
|
def close(self):
|
|
|
|
|
|
"""关闭数据库连接"""
|
|
|
|
|
|
self.driver.close()
|
2024-07-14 23:59:52 +08:00
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
def is_running(self):
|
|
|
|
|
|
"""检查图数据库是否正在运行"""
|
|
|
|
|
|
if not config.enable_knowledge_graph or not config.enable_knowledge_base:
|
|
|
|
|
|
return False
|
|
|
|
|
|
else:
|
|
|
|
|
|
return self.status == "open"
|
|
|
|
|
|
|
2024-09-06 12:54:17 +08:00
|
|
|
|
def get_sample_nodes(self, kgdb_name='neo4j', num=50):
|
2024-09-09 17:07:03 +08:00
|
|
|
|
"""获取指定数据库的 num 个节点信息"""
|
2024-09-06 12:54:17 +08:00
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
def query(tx, num):
|
|
|
|
|
|
result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(num))
|
|
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
return session.execute_read(query, num)
|
|
|
|
|
|
|
2024-07-15 18:58:34 +08:00
|
|
|
|
def create_graph_database(self, kgdb_name):
|
|
|
|
|
|
"""创建新的数据库,如果已存在则返回已有数据库的名称"""
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
existing_databases = session.run("SHOW DATABASES")
|
|
|
|
|
|
existing_db_names = [db['name'] for db in existing_databases]
|
2024-07-16 18:14:27 +08:00
|
|
|
|
|
2024-07-15 18:58:34 +08:00
|
|
|
|
if existing_db_names:
|
|
|
|
|
|
print(f"已存在数据库: {existing_db_names[0]}")
|
|
|
|
|
|
return existing_db_names[0] # 返回所有已有数据库名称
|
2024-07-16 18:14:27 +08:00
|
|
|
|
|
2024-07-15 18:58:34 +08:00
|
|
|
|
session.run(f"CREATE DATABASE {kgdb_name}")
|
|
|
|
|
|
print(f"数据库 '{kgdb_name}' 创建成功.")
|
|
|
|
|
|
return kgdb_name # 返回创建的数据库名称
|
2024-07-16 18:14:27 +08:00
|
|
|
|
|
2024-09-14 02:44:53 +08:00
|
|
|
|
def use_database(self, kgdb_name="neo4j"):
|
2024-07-15 18:58:34 +08:00
|
|
|
|
"""切换到指定数据库"""
|
2024-09-14 02:44:53 +08:00
|
|
|
|
assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致"
|
|
|
|
|
|
if self.status == "closed":
|
|
|
|
|
|
self.start()
|
2024-07-15 18:58:34 +08:00
|
|
|
|
|
2024-07-17 18:50:01 +08:00
|
|
|
|
def txt_add_entity(self, triples, kgdb_name='neo4j'):
|
2024-07-15 18:58:34 +08:00
|
|
|
|
"""添加实体三元组"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
def create(tx, triples):
|
|
|
|
|
|
for triple in triples:
|
|
|
|
|
|
h = triple['h']
|
|
|
|
|
|
t = triple['t']
|
|
|
|
|
|
r = triple['r']
|
|
|
|
|
|
query = (
|
|
|
|
|
|
"MERGE (a:Entity {name: $h}) "
|
|
|
|
|
|
"MERGE (b:Entity {name: $t}) "
|
|
|
|
|
|
"MERGE (a)-[:" + r.replace(" ", "_") + "]->(b)"
|
|
|
|
|
|
)
|
|
|
|
|
|
tx.run(query, h=h, t=t)
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
session.execute_write(create, triples)
|
|
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def txt_add_vector_entity(self, triples, kgdb_name='neo4j'):
|
|
|
|
|
|
"""添加实体三元组"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
def _index_exists(tx, index_name):
|
2025-02-28 00:38:59 +08:00
|
|
|
|
"""检查索引是否存在"""
|
2024-07-20 20:30:32 +08:00
|
|
|
|
result = tx.run("SHOW INDEXES")
|
|
|
|
|
|
for record in result:
|
|
|
|
|
|
if record["name"] == index_name:
|
|
|
|
|
|
return True
|
|
|
|
|
|
return False
|
2025-02-28 00:38:59 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def _create_graph(tx, data):
|
2025-02-28 00:38:59 +08:00
|
|
|
|
"""添加一个三元组"""
|
2024-07-20 20:30:32 +08:00
|
|
|
|
for entry in data:
|
|
|
|
|
|
tx.run("""
|
|
|
|
|
|
MERGE (h:Entity {name: $h})
|
|
|
|
|
|
MERGE (t:Entity {name: $t})
|
|
|
|
|
|
MERGE (h)-[r:RELATION {type: $r}]->(t)
|
|
|
|
|
|
""", h=entry['h'], t=entry['t'], r=entry['r'])
|
2025-02-28 00:38:59 +08:00
|
|
|
|
|
2024-09-06 12:54:17 +08:00
|
|
|
|
def _create_vector_index(tx, dim):
|
2025-02-28 00:38:59 +08:00
|
|
|
|
"""创建向量索引"""
|
|
|
|
|
|
# NOTE 这里是否是会重复构建索引?
|
2024-09-06 12:54:17 +08:00
|
|
|
|
index_name = "entityEmbeddings"
|
2024-07-20 20:30:32 +08:00
|
|
|
|
if not _index_exists(tx, index_name):
|
|
|
|
|
|
tx.run(f"""
|
|
|
|
|
|
CREATE VECTOR INDEX {index_name}
|
|
|
|
|
|
FOR (n: Entity) ON (n.embedding)
|
|
|
|
|
|
OPTIONS {{indexConfig: {{
|
2024-09-06 12:54:17 +08:00
|
|
|
|
`vector.dimensions`: {dim},
|
2024-07-20 20:30:32 +08:00
|
|
|
|
`vector.similarity_function`: 'cosine'
|
|
|
|
|
|
}} }};
|
|
|
|
|
|
""")
|
2024-09-06 12:54:17 +08:00
|
|
|
|
|
2025-02-28 00:38:59 +08:00
|
|
|
|
# 判断模型名称是否匹配
|
2024-09-06 12:54:17 +08:00
|
|
|
|
from src.config import EMBED_MODEL_INFO
|
2025-03-20 19:51:46 +08:00
|
|
|
|
cur_embed_info = EMBED_MODEL_INFO[config.embed_model]
|
2025-02-28 00:38:59 +08:00
|
|
|
|
self.embed_model_name = self.embed_model_name or cur_embed_info.get('name')
|
|
|
|
|
|
assert self.embed_model_name == cur_embed_info.get('name') or self.embed_model_name is None, \
|
|
|
|
|
|
f"embed_model_name={self.embed_model_name}, {cur_embed_info.get('name')=}"
|
|
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
with self.driver.session() as session:
|
2025-02-28 00:38:59 +08:00
|
|
|
|
logger.info(f"Adding entity to {kgdb_name}")
|
2024-07-20 20:30:32 +08:00
|
|
|
|
session.execute_write(_create_graph, triples)
|
2025-03-20 19:51:46 +08:00
|
|
|
|
logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}")
|
2025-02-28 00:38:59 +08:00
|
|
|
|
session.execute_write(_create_vector_index, cur_embed_info['dimension'])
|
|
|
|
|
|
# NOTE 这里需要异步处理
|
2024-09-06 12:54:17 +08:00
|
|
|
|
for i, entry in enumerate(triples):
|
2025-02-28 00:38:59 +08:00
|
|
|
|
logger.debug(f"Adding entity {i+1}/{len(triples)}")
|
2024-07-20 20:30:32 +08:00
|
|
|
|
embedding_h = self.get_embedding(entry['h'])
|
|
|
|
|
|
embedding_t = self.get_embedding(entry['t'])
|
2025-02-28 00:38:59 +08:00
|
|
|
|
session.execute_write(self.set_embedding, entry['h'], embedding_h)
|
2024-07-20 20:30:32 +08:00
|
|
|
|
session.execute_write(self.set_embedding, entry['t'], embedding_t)
|
|
|
|
|
|
|
2025-02-28 02:40:34 +08:00
|
|
|
|
# 数据添加完成后保存图信息
|
|
|
|
|
|
self.save_graph_info()
|
|
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'):
|
2024-07-24 18:32:14 +08:00
|
|
|
|
self.status = "processing"
|
2024-10-14 16:51:20 +08:00
|
|
|
|
kgdb_name = kgdb_name or 'neo4j'
|
2024-07-20 20:30:32 +08:00
|
|
|
|
self.use_database(kgdb_name) # 切换到指定数据库
|
2025-02-28 00:38:59 +08:00
|
|
|
|
logger.info(f"Start adding entity to {kgdb_name} with {file_path}")
|
2024-07-24 18:32:14 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def read_triples(file_path):
|
|
|
|
|
|
with open(file_path, 'r', encoding='utf-8') as file:
|
|
|
|
|
|
for line in file:
|
2024-07-24 19:35:33 +08:00
|
|
|
|
yield json.loads(line.strip())
|
|
|
|
|
|
|
|
|
|
|
|
triples = list(read_triples(file_path))
|
|
|
|
|
|
|
2024-09-06 12:54:17 +08:00
|
|
|
|
self.txt_add_vector_entity(triples, kgdb_name)
|
|
|
|
|
|
|
2024-07-24 18:32:14 +08:00
|
|
|
|
self.status = "open"
|
2025-02-28 02:40:34 +08:00
|
|
|
|
# 更新并保存图数据库信息
|
|
|
|
|
|
self.save_graph_info()
|
2024-07-20 20:30:32 +08:00
|
|
|
|
return kgdb_name
|
2024-07-14 23:59:52 +08:00
|
|
|
|
|
2024-07-15 18:58:34 +08:00
|
|
|
|
def delete_entity(self, entity_name=None, kgdb_name="neo4j"):
|
|
|
|
|
|
"""删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
if entity_name:
|
|
|
|
|
|
session.execute_write(self._delete_specific_entity, entity_name)
|
|
|
|
|
|
else:
|
|
|
|
|
|
session.execute_write(self._delete_all_entities)
|
|
|
|
|
|
|
|
|
|
|
|
def _delete_specific_entity(self, tx, entity_name):
|
|
|
|
|
|
query = """
|
|
|
|
|
|
MATCH (n {name: $entity_name})
|
|
|
|
|
|
DETACH DELETE n
|
|
|
|
|
|
"""
|
|
|
|
|
|
tx.run(query, entity_name=entity_name)
|
|
|
|
|
|
|
|
|
|
|
|
def _delete_all_entities(self, tx):
|
|
|
|
|
|
query = """
|
|
|
|
|
|
MATCH (n)
|
|
|
|
|
|
DETACH DELETE n
|
|
|
|
|
|
"""
|
|
|
|
|
|
tx.run(query)
|
|
|
|
|
|
|
2024-09-14 02:44:53 +08:00
|
|
|
|
def query_node(self, entity_name, hops=2, **kwargs):
|
|
|
|
|
|
# TODO 添加判断节点数量为 0 停止检索
|
2025-02-28 02:40:34 +08:00
|
|
|
|
|
|
|
|
|
|
logger.debug(f"Query graph node {entity_name} with {hops=}")
|
2024-09-14 02:44:53 +08:00
|
|
|
|
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):
|
2024-07-15 18:58:34 +08:00
|
|
|
|
self.use_database(kgdb_name)
|
2024-09-14 02:44:53 +08:00
|
|
|
|
def query(tx, text):
|
|
|
|
|
|
embedding = self.get_embedding(text)
|
|
|
|
|
|
result = tx.run("""
|
|
|
|
|
|
CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding)
|
|
|
|
|
|
YIELD node AS similarEntity, score
|
|
|
|
|
|
RETURN similarEntity.name AS name, score
|
|
|
|
|
|
""", embedding=embedding)
|
2024-07-15 18:58:34 +08:00
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
2025-02-28 00:38:59 +08:00
|
|
|
|
results = session.execute_read(query, entity_name)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# 筛选出分数高于阈值的实体
|
|
|
|
|
|
qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold]
|
|
|
|
|
|
logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}")
|
|
|
|
|
|
|
|
|
|
|
|
# 对每个合格的实体进行查询
|
|
|
|
|
|
all_query_results = []
|
|
|
|
|
|
for entity in qualified_entities:
|
|
|
|
|
|
query_result = self.query_specific_entity(entity_name=entity, hops=hops, kgdb_name=kgdb_name)
|
|
|
|
|
|
all_query_results.extend(query_result)
|
|
|
|
|
|
|
|
|
|
|
|
return all_query_results
|
2024-07-15 18:58:34 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2):
|
2025-02-28 00:38:59 +08:00
|
|
|
|
"""查询指定实体三元组信息(无向关系)"""
|
2024-07-15 18:58:34 +08:00
|
|
|
|
self.use_database(kgdb_name)
|
2024-07-16 12:27:46 +08:00
|
|
|
|
def query(tx, entity_name, hops):
|
|
|
|
|
|
result = tx.run(f"""
|
2025-02-28 00:38:59 +08:00
|
|
|
|
MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m)
|
2025-02-28 02:42:21 +08:00
|
|
|
|
RETURN n, r, m
|
2024-07-15 18:58:34 +08:00
|
|
|
|
""", entity_name=entity_name)
|
|
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
2024-07-17 19:34:05 +08:00
|
|
|
|
return session.execute_read(query, entity_name, hops)
|
2024-07-15 18:58:34 +08:00
|
|
|
|
|
2024-09-14 02:44:53 +08:00
|
|
|
|
def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2):
|
|
|
|
|
|
"""查询图数据库中所有三元组信息"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
def query(tx, hops):
|
|
|
|
|
|
result = tx.run(f"""
|
|
|
|
|
|
MATCH (n)-[r*1..{hops}]->(m)
|
|
|
|
|
|
RETURN n, r, m
|
|
|
|
|
|
""")
|
|
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
return session.execute_read(query, hops)
|
|
|
|
|
|
|
2024-07-17 18:50:01 +08:00
|
|
|
|
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
|
2024-07-15 18:58:34 +08:00
|
|
|
|
"""查询指定关系三元组信息"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
2024-07-16 12:27:46 +08:00
|
|
|
|
def query(tx, relationship_type, hops):
|
2024-07-15 18:58:34 +08:00
|
|
|
|
result = tx.run(f"""
|
2024-07-16 12:27:46 +08:00
|
|
|
|
MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m)
|
2024-07-15 18:58:34 +08:00
|
|
|
|
RETURN n, r, m
|
|
|
|
|
|
""")
|
|
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
2024-07-17 19:34:05 +08:00
|
|
|
|
return session.execute_read(query, relationship_type, hops)
|
2024-07-15 18:58:34 +08:00
|
|
|
|
|
2024-07-17 18:50:01 +08:00
|
|
|
|
def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2):
|
2024-07-15 18:58:34 +08:00
|
|
|
|
"""模糊查询"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
2024-07-16 12:27:46 +08:00
|
|
|
|
def query(tx, keyword, hops):
|
|
|
|
|
|
result = tx.run(f"""
|
2024-07-15 18:58:34 +08:00
|
|
|
|
MATCH (n:Entity)
|
|
|
|
|
|
WHERE n.name CONTAINS $keyword
|
2024-07-16 12:27:46 +08:00
|
|
|
|
MATCH (n)-[r*1..{hops}]->(m)
|
2024-07-15 18:58:34 +08:00
|
|
|
|
RETURN n, r, m
|
|
|
|
|
|
""", keyword=keyword)
|
|
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
2024-07-17 19:34:05 +08:00
|
|
|
|
return session.execute_read(query, keyword, hops)
|
2024-07-24 18:32:14 +08:00
|
|
|
|
|
2024-07-17 18:50:01 +08:00
|
|
|
|
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
|
2024-07-16 12:27:46 +08:00
|
|
|
|
"""查询指定节点的详细信息返回信息"""
|
2024-07-15 18:58:34 +08:00
|
|
|
|
self.use_database(kgdb_name) # 切换到指定数据库
|
2024-07-16 12:27:46 +08:00
|
|
|
|
def query(tx, node_name, hops):
|
|
|
|
|
|
result = tx.run(f"""
|
2024-07-16 18:14:27 +08:00
|
|
|
|
MATCH (n {{name: $node_name}})
|
|
|
|
|
|
OPTIONAL MATCH (n)-[r*1..{hops}]->(m)
|
2024-07-15 18:58:34 +08:00
|
|
|
|
RETURN n, r, m
|
|
|
|
|
|
""", node_name=node_name)
|
|
|
|
|
|
return result.values()
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
2024-07-17 19:34:05 +08:00
|
|
|
|
return session.execute_read(query, node_name, hops)
|
2024-07-24 18:32:14 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def get_embedding(self, text):
|
|
|
|
|
|
with torch.no_grad():
|
2025-03-20 19:51:46 +08:00
|
|
|
|
from src import knowledge_base
|
|
|
|
|
|
outputs = knowledge_base.embed_model.encode([text])[0]
|
2025-02-28 00:38:59 +08:00
|
|
|
|
return outputs
|
2024-07-24 18:32:14 +08:00
|
|
|
|
|
2024-07-20 20:30:32 +08:00
|
|
|
|
def set_embedding(self, tx, entity_name, embedding):
|
|
|
|
|
|
tx.run("""
|
|
|
|
|
|
MATCH (e:Entity {name: $name})
|
|
|
|
|
|
CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding)
|
|
|
|
|
|
""", name=entity_name, embedding=embedding)
|
2024-07-24 18:32:14 +08:00
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
def get_graph_info(self, graph_name="neo4j"):
|
|
|
|
|
|
self.use_database(graph_name)
|
|
|
|
|
|
def query(tx):
|
|
|
|
|
|
entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"]
|
|
|
|
|
|
relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"]
|
|
|
|
|
|
triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"]
|
|
|
|
|
|
|
|
|
|
|
|
# 获取所有标签
|
|
|
|
|
|
labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"]
|
|
|
|
|
|
|
|
|
|
|
|
return {
|
|
|
|
|
|
"graph_name": graph_name,
|
|
|
|
|
|
"entity_count": entity_count,
|
|
|
|
|
|
"relationship_count": relationship_count,
|
|
|
|
|
|
"triples_count": triples_count,
|
|
|
|
|
|
"labels": labels,
|
2025-02-28 02:40:34 +08:00
|
|
|
|
"status": self.status,
|
|
|
|
|
|
"embed_model_name": self.embed_model_name,
|
2025-03-20 19:51:46 +08:00
|
|
|
|
"unindexed_node_count": self.query_nodes_without_embedding(graph_name)
|
2025-02-28 02:40:34 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
2025-03-20 19:51:46 +08:00
|
|
|
|
try:
|
|
|
|
|
|
if self.status == "open" and self.driver and self.is_running():
|
|
|
|
|
|
# 获取数据库信息
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
graph_info = session.execute_read(query)
|
|
|
|
|
|
|
2025-02-28 02:40:34 +08:00
|
|
|
|
# 添加时间戳
|
|
|
|
|
|
from datetime import datetime
|
|
|
|
|
|
graph_info["last_updated"] = datetime.now().isoformat()
|
2025-03-20 19:51:46 +08:00
|
|
|
|
return graph_info
|
|
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"获取图数据库信息失败:{e}, {traceback.format_exc()}")
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
def save_graph_info(self, graph_name="neo4j"):
|
|
|
|
|
|
"""
|
|
|
|
|
|
将图数据库的基本信息保存到工作目录中的JSON文件
|
|
|
|
|
|
保存的信息包括:数据库名称、状态、嵌入模型名称等
|
|
|
|
|
|
"""
|
|
|
|
|
|
try:
|
|
|
|
|
|
graph_info = self.get_graph_info(graph_name)
|
|
|
|
|
|
if graph_info is None:
|
|
|
|
|
|
logger.error(f"图数据库信息为空,无法保存")
|
|
|
|
|
|
return False
|
2025-02-28 02:40:34 +08:00
|
|
|
|
|
|
|
|
|
|
info_file_path = os.path.join(self.work_dir, "graph_info.json")
|
|
|
|
|
|
with open(info_file_path, 'w', encoding='utf-8') as f:
|
|
|
|
|
|
json.dump(graph_info, f, ensure_ascii=False, indent=2)
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(f"图数据库信息已保存到:{info_file_path}")
|
|
|
|
|
|
return True
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"保存图数据库信息失败:{e}")
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
2025-03-08 22:49:13 +08:00
|
|
|
|
def query_nodes_without_embedding(self, kgdb_name='neo4j'):
|
|
|
|
|
|
"""查询没有嵌入向量的节点
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
list: 没有嵌入向量的节点列表
|
|
|
|
|
|
"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
|
|
|
|
|
|
def query(tx):
|
|
|
|
|
|
result = tx.run("""
|
|
|
|
|
|
MATCH (n:Entity)
|
|
|
|
|
|
WHERE n.embedding IS NULL
|
|
|
|
|
|
RETURN n.name AS name
|
|
|
|
|
|
""")
|
|
|
|
|
|
return [record["name"] for record in result]
|
|
|
|
|
|
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
return session.execute_read(query)
|
|
|
|
|
|
|
2025-02-28 02:40:34 +08:00
|
|
|
|
def load_graph_info(self):
|
|
|
|
|
|
"""
|
|
|
|
|
|
从工作目录中的JSON文件加载图数据库的基本信息
|
|
|
|
|
|
返回True表示加载成功,False表示加载失败
|
|
|
|
|
|
"""
|
|
|
|
|
|
try:
|
|
|
|
|
|
info_file_path = os.path.join(self.work_dir, "graph_info.json")
|
|
|
|
|
|
if not os.path.exists(info_file_path):
|
|
|
|
|
|
logger.warning(f"图数据库信息文件不存在:{info_file_path}")
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
with open(info_file_path, 'r', encoding='utf-8') as f:
|
|
|
|
|
|
graph_info = json.load(f)
|
|
|
|
|
|
|
|
|
|
|
|
# 更新对象属性
|
|
|
|
|
|
if graph_info.get("embed_model_name"):
|
|
|
|
|
|
self.embed_model_name = graph_info["embed_model_name"]
|
|
|
|
|
|
|
|
|
|
|
|
# 如果需要,可以加载更多信息
|
|
|
|
|
|
# 注意:这里不更新self.kgdb_name,因为它是在初始化时设置的
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(f"已加载图数据库信息,最后更新时间:{graph_info.get('last_updated')}")
|
|
|
|
|
|
return True
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"加载图数据库信息失败:{e}")
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
2025-03-08 22:49:13 +08:00
|
|
|
|
def add_embedding_to_nodes(self, node_names=None, kgdb_name='neo4j'):
|
|
|
|
|
|
"""为节点添加嵌入向量
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
node_names (list, optional): 要添加嵌入向量的节点名称列表,None表示所有没有嵌入向量的节点
|
|
|
|
|
|
kgdb_name (str, optional): 图数据库名称,默认为'neo4j'
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
int: 成功添加嵌入向量的节点数量
|
|
|
|
|
|
"""
|
|
|
|
|
|
self.use_database(kgdb_name)
|
|
|
|
|
|
|
|
|
|
|
|
# 如果node_names为None,则获取所有没有嵌入向量的节点
|
|
|
|
|
|
if node_names is None:
|
|
|
|
|
|
node_names = self.query_nodes_without_embedding(kgdb_name)
|
|
|
|
|
|
|
|
|
|
|
|
count = 0
|
|
|
|
|
|
with self.driver.session() as session:
|
|
|
|
|
|
for node_name in node_names:
|
|
|
|
|
|
try:
|
|
|
|
|
|
embedding = self.get_embedding(node_name)
|
|
|
|
|
|
session.execute_write(self.set_embedding, node_name, embedding)
|
|
|
|
|
|
count += 1
|
|
|
|
|
|
except Exception as e:
|
2025-03-19 01:21:51 +08:00
|
|
|
|
logger.error(f"为节点 '{node_name}' 添加嵌入向量失败: {e}, {traceback.format_exc()}")
|
2025-03-08 22:49:13 +08:00
|
|
|
|
|
|
|
|
|
|
return count
|
|
|
|
|
|
|
2024-07-15 18:58:34 +08:00
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
2025-02-28 00:38:59 +08:00
|
|
|
|
pass
|