暂存
This commit is contained in:
parent
582e367f52
commit
99c7ac659b
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
@ -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}
|
||||
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user