ForcePilot/src/core/graphbase.py
2024-09-18 21:03:12 +08:00

537 lines
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
import json
import torch
from neo4j import GraphDatabase as GD
# from src.plugins import pdf2txt, OneKE
from transformers import AutoTokenizer, AutoModel
from FlagEmbedding import FlagModel, FlagReranker
import warnings
from src.plugins import pdf2txt
from src.plugins.oneke import OneKE
from src.utils import setup_logger
logger = setup_logger("server-graphbase")
warnings.filterwarnings("ignore", category=UserWarning)
"""
from neo4j import GraphDatabase
import random
class KnowledgeGraph:
def __init__(self, uri, user, password):
self._driver = GraphDatabase.driver(uri, auth=(user, password))
def close(self):
self._driver.close()
def use_database(self, kgdb_name):
with self._driver.session() as session:
session.run(f"USE {kgdb_name}")
def get_sample_nodes(self, kgdb_name='neo4j', num=50):
self.use_database(kgdb_name)
selected_nodes = set()
nodes_to_expand = set()
result_nodes = []
with self._driver.session() as session:
while len(selected_nodes) < num:
# 如果需要扩展的节点为空,随机选择一个新节点
if not nodes_to_expand:
result = session.run("MATCH (n) RETURN n, rand() as r ORDER BY r LIMIT 1")
for record in result:
nodes_to_expand.add(record['n'].id)
result_nodes.append({'n': record['n'], 'r': None, 'm': None})
# 从需要扩展的节点中随机选择一个节点
current_node_id = random.choice(list(nodes_to_expand))
nodes_to_expand.remove(current_node_id)
# 获取当前节点的邻居节点最多5个
result = session.run(
f"MATCH (n)-[r]-(m) WHERE id(n) = {current_node_id} RETURN n, r, m LIMIT 5"
)
for record in result:
neighbor_node_id = record['m'].id
if neighbor_node_id not in selected_nodes:
selected_nodes.add(neighbor_node_id)
nodes_to_expand.add(neighbor_node_id)
result_nodes.append({'n': record['n'], 'r': record['r'], 'm': record['m']})
# 如果已经达到最大值,停止扩展
if len(selected_nodes) >= num:
break
return result_nodes[:num]
# 示例用法
uri = "bolt://localhost:7687"
user = "neo4j"
password = "password"
kg = KnowledgeGraph(uri, user, password)
sample_nodes = kg.get_sample_nodes(num=50)
for node in sample_nodes:
print(f"Node: {node['n'].id}, Relationship: {node['r']}, Neighbor: {node['m'].id}")
kg.close()
"""
UIE_MODEL = None
class GraphDatabase:
def __init__(self, config, embed_model=None, kgdb_name="neo4j"):
self.config = config
self.driver = None
self.files = []
self.status = "closed"
self.kgdb_name = kgdb_name
assert embed_model, "embed_model=None"
self.embed_model = embed_model
def start(self):
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}")
self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password))
self.status = "open"
def close(self):
"""关闭数据库连接"""
self.driver.close()
def get_sample_nodes(self, kgdb_name='neo4j', num=50):
"""获取指定数据库的 num 个节点信息"""
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)
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]
if existing_db_names:
print(f"已存在数据库: {existing_db_names[0]}")
return existing_db_names[0] # 返回所有已有数据库名称
session.run(f"CREATE DATABASE {kgdb_name}")
print(f"数据库 '{kgdb_name}' 创建成功.")
return kgdb_name # 返回创建的数据库名称
def get_database_info(self, db_name="neo4j"):
"""获取指定数据库的信息"""
self.use_database(db_name)
def query(tx):
entity_count = tx.run("MATCH (n:Entity) 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"]
return {
"database_name": db_name,
"entity_count": entity_count,
"relationship_count": relationship_count,
"triples_count": triples_count,
"status": self.status
}
with self.driver.session() as session:
return session.execute_read(query)
def use_database(self, kgdb_name="neo4j"):
"""切换到指定数据库"""
assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致"
if self.status == "closed":
self.start()
def txt_add_entity(self, triples, kgdb_name='neo4j'):
"""添加实体三元组"""
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)
def pdf_file_add_entity(self, file_path, output_path, kgdb_name='neo4j'):
self.use_database(kgdb_name) # 切换到指定数据库
text_path = pdf2txt(file_path)
global UIE_MODEL
if UIE_MODEL is None:
UIE_MODEL = OneKE()
triples_path = UIE_MODEL.processing_text_to_kg(text_path, output_path)
self.jsonl_file_add_entity(triples_path)
return kgdb_name
def txt_add_vector_entity(self, triples, kgdb_name='neo4j'):
"""添加实体三元组"""
self.use_database(kgdb_name)
def _index_exists(tx, index_name):
result = tx.run("SHOW INDEXES")
for record in result:
if record["name"] == index_name:
return True
return False
def _create_graph(tx, data):
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'])
def _create_vector_index(tx, dim):
index_name = "entityEmbeddings"
if not _index_exists(tx, index_name):
tx.run(f"""
CREATE VECTOR INDEX {index_name}
FOR (n: Entity) ON (n.embedding)
OPTIONS {{indexConfig: {{
`vector.dimensions`: {dim},
`vector.similarity_function`: 'cosine'
}} }};
""")
from src.config import EMBED_MODEL_INFO
embed_info = EMBED_MODEL_INFO[self.config.embed_model]
with self.driver.session() as session:
session.execute_write(_create_graph, triples)
session.execute_write(_create_vector_index, embed_info.dimension)
for i, entry in enumerate(triples):
logger.info(f"Adding entity {i+1}/{len(triples)}")
embedding_h = self.get_embedding(entry['h'])
session.execute_write(self.set_embedding, entry['h'], embedding_h)
embedding_t = self.get_embedding(entry['t'])
session.execute_write(self.set_embedding, entry['t'], embedding_t)
def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'):
self.status = "processing"
self.use_database(kgdb_name) # 切换到指定数据库
def read_triples(file_path):
with open(file_path, 'r', encoding='utf-8') as file:
for line in file:
yield json.loads(line.strip())
triples = list(read_triples(file_path))
self.txt_add_vector_entity(triples, kgdb_name)
self.status = "open"
return kgdb_name
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)
def query_node(self, entity_name, hops=2, **kwargs):
# TODO 添加判断节点数量为 0 停止检索
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):
results = self.query_by_vector_tep(entity_name=entity_name)
# 筛选出分数高于阈值的实体
qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold]
# 对每个合格的实体进行查询
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
def query_by_vector_tep(self, entity_name, kgdb_name='neo4j'):
"""向量查询"""
self.use_database(kgdb_name)
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)
return result.values()
with self.driver.session() as session:
return session.execute_read(query, entity_name)
def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2):
"""查询指定实体三元组信息"""
self.use_database(kgdb_name)
def query(tx, entity_name, hops):
result = tx.run(f"""
MATCH (n {{name: $entity_name}})-[r*1..{hops}]->(m)
RETURN n.name AS node_name, r, m.name AS neighbor_name
""", entity_name=entity_name)
return result.values()
with self.driver.session() as session:
return session.execute_read(query, entity_name, hops)
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)
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
"""查询指定关系三元组信息"""
self.use_database(kgdb_name)
def query(tx, relationship_type, hops):
result = tx.run(f"""
MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m)
RETURN n, r, m
""")
return result.values()
with self.driver.session() as session:
return session.execute_read(query, relationship_type, hops)
def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2):
"""模糊查询"""
self.use_database(kgdb_name)
def query(tx, keyword, hops):
result = tx.run(f"""
MATCH (n:Entity)
WHERE n.name CONTAINS $keyword
MATCH (n)-[r*1..{hops}]->(m)
RETURN n, r, m
""", keyword=keyword)
return result.values()
with self.driver.session() as session:
return session.execute_read(query, keyword, hops)
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
"""查询指定节点的详细信息返回信息"""
self.use_database(kgdb_name) # 切换到指定数据库
def query(tx, node_name, hops):
result = tx.run(f"""
MATCH (n {{name: $node_name}})
OPTIONAL MATCH (n)-[r*1..{hops}]->(m)
RETURN n, r, m
""", node_name=node_name)
return result.values()
with self.driver.session() as session:
return session.execute_read(query, node_name, hops)
def get_embedding(self, text):
inputs = [text]
with torch.no_grad():
outputs = self.embed_model.encode(inputs)
embeddings = outputs[0] # 假设取平均作为文本的嵌入向量
return embeddings
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)
# def format_query_results(self, results):
# formatted_results = []
# for row in results:
# n, rs, m = row
# entity_a = n['name']
# entity_b = m['name']
# for rel in rs:
# relationship = rel.type
# formatted_results.append(f"实体 {entity_a} 和 实体 {entity_b} 的关系是 {relationship}")
# return formatted_results
if __name__ == "__main__":
config = None
kgdb_name = "neo4j"
class EmbeddingModel(FlagModel):
def __init__(self, config, **kwargs):
model_name_or_path = "/data2024/yyyl/model/BAAI/bge-large-zh-v1.5/"
super().__init__(model_name_or_path, use_fp16=False, **kwargs)
model = EmbeddingModel(config)
# 初始化知识图谱数据库
kgdb = GraphDatabase(config, model)
# 创建新的数据库
# kgdb.create_graph_database("db2")
# 返回指定数据库信息
# info = kgdb.get_database_info(kgdb_name)
# print(info)
# triples = [
# {
# "h": "CCC",
# "t": "EE",
# "r": "同学"
# },
# {
# "h": "EE",
# "t": "RR",
# "r": "同事"
# }
# ]
# kgdb.query_by_vector("z")
# def format_query_results(results):
# formatted_results = {"nodes": [], "edges": []}
# # 用于存储所有唯一的节点信息
# node_dict = {}
# for item in results:
# # 确保item[1]是一个非空的列表
# if isinstance(item[1], list) and len(item[1]) > 0:
# relationship = item[1][0]
# rel_id = relationship.element_id
# nodes = relationship.nodes
# if len(nodes) == 2:
# node1, node2 = nodes
# # 提取源节点和目标节点信息
# node1_id = node1.element_id
# node2_id = node2.element_id
# node1_name = item[0] # 假设节点名称和列表中的第一个元素相同
# node2_name = item[2] if len(item) > 2 else 'unknown'
# # 记录节点信息
# if node1_id not in node_dict:
# node_dict[node1_id] = {"id": node1_id, "name": node1_name}
# if node2_id not in node_dict:
# node_dict[node2_id] = {"id": node2_id, "name": node2_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": node1_id,
# "target_id": node2_id,
# "source_name": node1_name,
# "target_name": node2_name
# })
# # 将唯一的节点信息添加到结果中
# formatted_results["nodes"] = list(node_dict.values())
# return formatted_results
# entities = ['jqy', '维c']
# results = []
# for entitie in entities:
# result = kgdb.query_by_vector(entitie)
# if result != []:
# results.extend(result)
# print(format_query_results(results))
# kgdb.txt_add_vector_entity(triples, model)
# print("Extend the Graph data base")
# kgdb.jsonl_file_add_entity("/data2024/yyyl/ProjectAthena/tep.jsonl", kgdb_name)
# print("Extend the Graph data base")
# triples_path = "output.jsonl"
# def read_triples(file_path):
# with open(file_path, 'r', encoding='utf-8') as file:
# for line in file:
# item = json.loads(line.strip())
# yield [item]
# for trio in read_triples(triples_path):
# kgdb.txt_add_entity(trio)
# 通过文件添加三元组数据
# kgdb.file_add_entity("/data2024/yyyl/ProjectAthena/test.pdf", "/data2024/yyyl/ProjectAthena/output.jsonl", kgdb_name)
# print("Extend the Graph data base")
# 删除数据库信息
# kgdb.delete_entity()
# print("Clear the Graph data base")
# 查询所有节点和关系
# results = kgdb.query_all_nodes_and_relationships(kgdb_name)
# print(results)
# 查询特定实体及其关系
# results = kgdb.query_specific_entity("kllll", kgdb_name)
# print(results)
# 查询特定实体及其关系
# results = kgdb.query_by_vector_tep("z", model, tokenizer)
# print(results)
# 查询特定关系类型的所有节点
# results = kgdb.query_by_relationship_type("作用或食用效果", kgdb_name)
# print(results)
# 模糊查询
# results = kgdb.query_entity_like("三七", kgdb_name)
# print(results)
# 查询节点信息
# results = kgdb.query_entity_like("三七提取物", kgdb_name)
# print(results)
# 关闭数据库连接
kgdb.close()