diff --git a/.gitignore b/.gitignore index 27d9101c..90f3cb2c 100644 --- a/.gitignore +++ b/.gitignore @@ -29,4 +29,5 @@ cache *.pdf src/data -neo4j* \ No newline at end of file +neo4j* +*/package-lock.json \ No newline at end of file diff --git a/src/core/database.py b/src/core/database.py index 6810b0ff..baefefbf 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -6,6 +6,7 @@ from plugins import pdf2txt from core.knowledgebase import KnowledgeBase from core.filereader import pdfreader, plainreader from core.graphbase import GraphDatabase +from models.embedding import EmbeddingModel logger = setup_logger("DataBaseManager") @@ -20,6 +21,7 @@ class DataBaseLite: self.metadata = kwargs.get("metaname", {}) self.files = kwargs.get("files", []) + def update(self, metadata): self.metadata = metadata @@ -45,11 +47,12 @@ class DataBaseManager: def __init__(self, config=None) -> None: self.config = config self.database_path = "data/databases.json" - self.knowledge_base = KnowledgeBase(config) + self.embed_model = EmbeddingModel(config) + self.knowledge_base = KnowledgeBase(config, self.embed_model) self.data = {"databases": [], "graph": {}} if self.config.enable_knowledge_graph: - self.graph_base = GraphDatabase(self.config) + self.graph_base = GraphDatabase(self.config, self.embed_model) self.graph_base.start() self._load_databases() diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 12beacf9..0d745248 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -1,14 +1,21 @@ import os import json +import torch from neo4j import GraphDatabase as GD -from plugins import pdf2txt, OneKE +# from plugins import pdf2txt, OneKE +from transformers import AutoTokenizer, AutoModel +from FlagEmbedding import FlagModel, FlagReranker + + class GraphDatabase: - def __init__(self, config): + def __init__(self, config, embed_model=None): self.config = config self.driver = None self.files = [] self.status = "closed" + assert embed_model, "embed_model=None" + self.embed_model = embed_model def start(self): uri = os.environ.get("NEO4J_URI") @@ -75,7 +82,7 @@ class GraphDatabase: with self.driver.session() as session: session.execute_write(create, triples) - def file_add_entity(self, file_path, output_path, kgdb_name='neo4j'): + def pdf_file_add_entity(self, file_path, output_path, kgdb_name='neo4j'): self.use_database(kgdb_name) # 切换到指定数据库 text_path = pdf2txt(file_path) oneke = OneKE() @@ -87,7 +94,56 @@ class GraphDatabase: yield [item] for trio in read_triples(triples_path): self.txt_add_entity(trio, kgdb_name) - pass + 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): + index_name = "entity-embeddings" + 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`: 1024, + `vector.similarity_function`: 'cosine' + }} }}; + """) + with self.driver.session() as session: + session.execute_write(_create_graph, triples) + session.execute_write(_create_vector_index) + for entry in 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.use_database(kgdb_name) # 切换到指定数据库 + triples_path = file_path + 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): + self.txt_add_entity(trio, kgdb_name) + return kgdb_name def delete_entity(self, entity_name=None, kgdb_name="neo4j"): """删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体""" @@ -125,13 +181,13 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query, hops) - def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops = 2): + 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, r, m + RETURN n.name AS node_name, r, m.name AS neighbor_name """, entity_name=entity_name) return result.values() @@ -165,6 +221,29 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query, keyword, hops) + + def query_by_vector_tep(self, keyword, 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('entity-embeddings', 10, $embedding) + YIELD node AS similarEntity, score + RETURN similarEntity.name AS name, score + """, embedding=embedding) + # result = result.values() + # query = result[0][0] + return result.values() + + with self.driver.session() as session: + return session.execute_read(query, keyword) + + def query_by_vector(self, entity_name, kgdb_name='neo4j'): + self.use_database(kgdb_name) + result = self.query_by_vector_tep(entity_name) + ans = self.query_specific_entity(result[0][0]) + return ans def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): """查询指定节点的详细信息返回信息""" @@ -180,6 +259,19 @@ class GraphDatabase: 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: @@ -195,48 +287,117 @@ class GraphDatabase: if __name__ == "__main__": config = None - # 初始化知识图谱数据库 - kgdb = GraphDatabase(config) + kgdb_name = "neo4j" + QUERY_INSTRUCTION = { + "bge-large-zh-v1.5": "为这个句子生成表示以用于检索相关文章:", + } + 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": "莲子", - # "t": "生吃微甜,一煮就酥,食之软糯清香", - # "r": "用途" + # "h": "CCC", + # "t": "EE", + # "r": "同学" # }, # { - # "h": "食用菌类", - # "t": "蔬菜", - # "r": "用途" - # }, - # { - # "h": "野生蕈", - # "t": "食用菌", - # "r": "作用或食用效果" - # }, - # { - # "h": "大白菜", - # "t": "选择耐贮的晚熟品种,如小青口、核桃纹、抱头青、拧心青等", - # "r": "贮存方法" - # }, - # { - # "h": "大白菜", - # "t": "刚买回来的白菜,水分大,须晾晒三五天,白菜外叶失去部分水分发时,再撕去黄叶,堆码", - # "r": "贮存方法" + # "h": "EE", + # "t": "RR", + # "r": "同事" # } # ] + + # print(kgdb.query_by_vector("维C")) + 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.txt_add_entity(triples, 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") @@ -250,7 +411,11 @@ if __name__ == "__main__": # print(results) # 查询特定实体及其关系 - # results = kgdb.query_specific_entity("三七提取物", kgdb_name) + # results = kgdb.query_specific_entity("zzz", kgdb_name) + # print(results) + + # 查询特定实体及其关系 + # results = kgdb.query_by_vector_tep("z", model, tokenizer) # print(results) # 查询特定关系类型的所有节点 @@ -267,3 +432,4 @@ if __name__ == "__main__": # 关闭数据库连接 kgdb.close() + diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index f7a3165b..fd4946eb 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -8,11 +8,11 @@ logger = setup_logger("KnowledgeBase") class KnowledgeBase: - def __init__(self, config=None) -> None: + def __init__(self, config=None, embed_model=None ) -> None: self.config = config self._init_config(config) - - self.embed_model = EmbeddingModel(config) + assert embed_model, "embed_model=None" + self.embed_model = embed_model self.client = MilvusClient("data/vector_base/milvus.db") def _init_config(self, config): diff --git a/src/core/retriever.py b/src/core/retriever.py index 524c549c..d23ad00f 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,5 +1,8 @@ from core.startup import dbm, model from models.embedding import ReRanker +from utils.logging_config import setup_logger +logger = setup_logger("server-common") + class Retriever: @@ -30,10 +33,10 @@ class Retriever: kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res]) external += f"知识库信息: \n\n{kb_text}" - # db_res = refs.get("graph_base").get("results", []) - # if len(db_res) > 0: - # db_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in db_res]) - # external += f"图数据库信息: \n\n{db_text}" + db_res = refs.get("graph_base").get("results", []) + if len(db_res) > 0: + db_text = '\n'.join([f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res['edges']]) + external += f"图数据库信息: \n\n{db_text}" if len(external) > 0: query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}" @@ -53,9 +56,9 @@ class Retriever: results = [] if meta.get("use_graph"): for entitie in entities: - result = dbm.graph_base.query_entity_like(entitie) - results.extend(result) if result else None - + result = dbm.graph_base.query_by_vector(entitie) + if result != []: + results.extend(result) return {"results": self.format_query_results(results)} def query_knowledgebase(self, query, history, meta): @@ -113,29 +116,44 @@ class Retriever: return rewritten_query, entities - def format_query_results(self, results): + def format_query_results(sfle, results): formatted_results = {"nodes": [], "edges": []} - for row in results: - n, relations, m = row - formatted_results["nodes"].append({ - "id": n.id, - "name": n._properties["name"], - "properties": n._properties - }) - formatted_results["nodes"].append({ - "id": m.id, - "name": m._properties["name"], - "properties": m._properties - }) - for rel in relations: - formatted_results["edges"].append({ - "id": rel.id, - "type": rel.type, - "source": rel.start_node.id, - "target": rel.end_node.id, - "source_name": rel.start_node._properties["name"], - "target_name": rel.end_node._properties["name"], - }) + + node_dict = {} + + for item in results: + 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 def __call__(self, query, history, meta):