diff --git a/docker/web.Dockerfile b/docker/web.Dockerfile index 20737eb3..3a3624fb 100644 --- a/docker/web.Dockerfile +++ b/docker/web.Dockerfile @@ -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 diff --git a/src/core/database.py b/src/core/database.py index 99ee2832..2b80e22c 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -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 diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 399ec448..8bf895d6 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -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 \ No newline at end of file diff --git a/src/routers/base_router.py b/src/routers/base_router.py index b89eff59..0dcbbde7 100644 --- a/src/routers/base_router.py +++ b/src/routers/base_router.py @@ -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}