This commit is contained in:
Wenjie Zhang 2025-02-28 00:38:59 +08:00
parent 582e367f52
commit 99c7ac659b
4 changed files with 60 additions and 205 deletions

View File

@ -6,8 +6,8 @@ WORKDIR /app
COPY ./web/package*.json ./
# 安装依赖
RUN npm install --verbose --force
# RUN npm install --registry http://mirrors.cloud.tencent.com/npm/ --verbose --force
# RUN npm install --verbose --force
RUN npm install --registry http://mirrors.cloud.tencent.com/npm/ --verbose --force
# 复制源代码
COPY ./web .
@ -22,8 +22,8 @@ FROM node:latest AS build-stage
WORKDIR /app
COPY ./web/package*.json ./
RUN npm install --force
# RUN npm install --registry https://registry.npmmirror.com --force
# RUN npm install --force
RUN npm install --registry https://registry.npmmirror.com --force
COPY ./web .
RUN npm run build

View File

@ -19,7 +19,6 @@ class DataBaseManager:
if self.config.enable_knowledge_graph:
from src.core.graphbase import GraphDatabase
self.graph_base = GraphDatabase(self.config, self.embed_model)
self.graph_base.start()
else:
self.graph_base = None

View File

@ -9,9 +9,7 @@ import warnings
from src.plugins import pdf2txt
from src.plugins.oneke import OneKE
from src.utils import setup_logger
logger = setup_logger("server-graphbase")
from src.utils import logger
warnings.filterwarnings("ignore", category=UserWarning)
@ -28,6 +26,10 @@ class GraphDatabase:
assert embed_model, "embed_model=None"
self.embed_model = embed_model
self.embed_model_name = None
self.work_dir = os.path.join(config.save_dir, "knowledge_graph", kgdb_name)
os.makedirs(self.work_dir, exist_ok=True)
self.start()
def start(self):
uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687")
@ -42,7 +44,6 @@ class GraphDatabase:
logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}")
self.config.enable_knowledge_graph = False
def close(self):
"""关闭数据库连接"""
self.driver.close()
@ -119,33 +120,39 @@ class GraphDatabase:
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 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):
"""创建向量索引"""
# NOTE 这里是否是会重复构建索引?
index_name = "entityEmbeddings"
if not _index_exists(tx, index_name):
tx.run(f"""
@ -157,24 +164,31 @@ class GraphDatabase:
}} }};
""")
# 判断模型名称是否匹配
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.get('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)
cur_embed_info = EMBED_MODEL_INFO[self.config.embed_model]
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')=}"
with self.driver.session() as session:
logger.info(f"Adding entity to {kgdb_name}")
session.execute_write(_create_graph, triples)
logger.info(f"Creating vector index for {kgdb_name} with {self.config.embed_model}")
session.execute_write(_create_vector_index, cur_embed_info['dimension'])
# NOTE 这里需要异步处理
for i, entry in enumerate(triples):
logger.debug(f"Adding entity {i+1}/{len(triples)}")
embedding_h = self.get_embedding(entry['h'])
embedding_t = self.get_embedding(entry['t'])
session.execute_write(self.set_embedding, entry['h'], embedding_h)
session.execute_write(self.set_embedding, entry['t'], embedding_t)
def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'):
self.status = "processing"
kgdb_name = kgdb_name or 'neo4j'
self.embed_model_name = self.embed_model_name or self.config.embed_model
self.use_database(kgdb_name) # 切换到指定数据库
logger.info(f"Start adding entity to {kgdb_name} with {file_path}")
def read_triples(file_path):
with open(file_path, 'r', encoding='utf-8') as file:
@ -219,21 +233,6 @@ class GraphDatabase:
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)
@ -245,14 +244,27 @@ class GraphDatabase:
return result.values()
with self.driver.session() as session:
return session.execute_read(query, entity_name)
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
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)
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()
@ -316,11 +328,9 @@ class GraphDatabase:
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
outputs = self.embed_model.encode([text])[0]
return outputs
def set_embedding(self, tx, entity_name, embedding):
tx.run("""
@ -328,160 +338,6 @@ class GraphDatabase:
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()
pass

View File

@ -39,6 +39,6 @@ def get_log():
last_lines = deque(f, maxlen=1000)
log = ''.join(last_lines)
return {"log": log}
return {"log": log, "message": "success", "log_file": LOG_FILE}