知识图谱检索问答

This commit is contained in:
QiyiJiang 2024-07-20 12:30:32 +00:00
parent 789ab07f51
commit 956a8d1b02
5 changed files with 255 additions and 67 deletions

1
.gitignore vendored
View File

@ -30,3 +30,4 @@ cache
*.pdf *.pdf
src/data src/data
neo4j* neo4j*
*/package-lock.json

View File

@ -6,6 +6,7 @@ from plugins import pdf2txt
from core.knowledgebase import KnowledgeBase from core.knowledgebase import KnowledgeBase
from core.filereader import pdfreader, plainreader from core.filereader import pdfreader, plainreader
from core.graphbase import GraphDatabase from core.graphbase import GraphDatabase
from models.embedding import EmbeddingModel
logger = setup_logger("DataBaseManager") logger = setup_logger("DataBaseManager")
@ -20,6 +21,7 @@ class DataBaseLite:
self.metadata = kwargs.get("metaname", {}) self.metadata = kwargs.get("metaname", {})
self.files = kwargs.get("files", []) self.files = kwargs.get("files", [])
def update(self, metadata): def update(self, metadata):
self.metadata = metadata self.metadata = metadata
@ -45,11 +47,12 @@ class DataBaseManager:
def __init__(self, config=None) -> None: def __init__(self, config=None) -> None:
self.config = config self.config = config
self.database_path = "data/databases.json" 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": {}} self.data = {"databases": [], "graph": {}}
if self.config.enable_knowledge_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.graph_base.start()
self._load_databases() self._load_databases()

View File

@ -1,14 +1,21 @@
import os import os
import json import json
import torch
from neo4j import GraphDatabase as GD 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: class GraphDatabase:
def __init__(self, config): def __init__(self, config, embed_model=None):
self.config = config self.config = config
self.driver = None self.driver = None
self.files = [] self.files = []
self.status = "closed" self.status = "closed"
assert embed_model, "embed_model=None"
self.embed_model = embed_model
def start(self): def start(self):
uri = os.environ.get("NEO4J_URI") uri = os.environ.get("NEO4J_URI")
@ -75,7 +82,7 @@ class GraphDatabase:
with self.driver.session() as session: with self.driver.session() as session:
session.execute_write(create, triples) 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) # 切换到指定数据库 self.use_database(kgdb_name) # 切换到指定数据库
text_path = pdf2txt(file_path) text_path = pdf2txt(file_path)
oneke = OneKE() oneke = OneKE()
@ -87,7 +94,56 @@ class GraphDatabase:
yield [item] yield [item]
for trio in read_triples(triples_path): for trio in read_triples(triples_path):
self.txt_add_entity(trio, kgdb_name) 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"): def delete_entity(self, entity_name=None, kgdb_name="neo4j"):
"""删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体""" """删除数据库中的指定实体三元组, 参数entity_name为空则删除全部实体"""
@ -131,7 +187,7 @@ class GraphDatabase:
def query(tx, entity_name, hops): def query(tx, entity_name, hops):
result = tx.run(f""" result = tx.run(f"""
MATCH (n {{name: $entity_name}})-[r*1..{hops}]->(m) 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) """, entity_name=entity_name)
return result.values() return result.values()
@ -166,6 +222,29 @@ class GraphDatabase:
with self.driver.session() as session: with self.driver.session() as session:
return session.execute_read(query, keyword, hops) 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): def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
"""查询指定节点的详细信息返回信息""" """查询指定节点的详细信息返回信息"""
self.use_database(kgdb_name) # 切换到指定数据库 self.use_database(kgdb_name) # 切换到指定数据库
@ -180,6 +259,19 @@ class GraphDatabase:
with self.driver.session() as session: with self.driver.session() as session:
return session.execute_read(query, node_name, hops) 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): # def format_query_results(self, results):
# formatted_results = [] # formatted_results = []
# for row in results: # for row in results:
@ -195,10 +287,24 @@ class GraphDatabase:
if __name__ == "__main__": if __name__ == "__main__":
config = None config = None
# 初始化知识图谱数据库
kgdb = GraphDatabase(config)
kgdb_name = "neo4j" 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") # kgdb.create_graph_database("db2")
@ -206,37 +312,92 @@ if __name__ == "__main__":
# info = kgdb.get_database_info(kgdb_name) # info = kgdb.get_database_info(kgdb_name)
# print(info) # print(info)
# 通过文本添加三元组数据
# triples = [ # triples = [
# { # {
# "h": "莲子", # "h": "CCC",
# "t": "生吃微甜,一煮就酥,食之软糯清香", # "t": "EE",
# "r": "用途" # "r": "同学"
# }, # },
# { # {
# "h": "食用菌类", # "h": "EE",
# "t": "蔬菜", # "t": "RR",
# "r": "用途" # "r": "同事"
# },
# {
# "h": "野生蕈",
# "t": "食用菌",
# "r": "作用或食用效果"
# },
# {
# "h": "大白菜",
# "t": "选择耐贮的晚熟品种,如小青口、核桃纹、抱头青、拧心青等",
# "r": "贮存方法"
# },
# {
# "h": "大白菜",
# "t": "刚买回来的白菜,水分大,须晾晒三五天,白菜外叶失去部分水分发时,再撕去黄叶,堆码",
# "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) # kgdb.txt_add_entity(triples, kgdb_name)
# print("Extend the Graph data base") # 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) # kgdb.file_add_entity("/data2024/yyyl/ProjectAthena/test.pdf", "/data2024/yyyl/ProjectAthena/output.jsonl", kgdb_name)
# print("Extend the Graph data base") # print("Extend the Graph data base")
@ -250,7 +411,11 @@ if __name__ == "__main__":
# print(results) # 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) # print(results)
# 查询特定关系类型的所有节点 # 查询特定关系类型的所有节点
@ -267,3 +432,4 @@ if __name__ == "__main__":
# 关闭数据库连接 # 关闭数据库连接
kgdb.close() kgdb.close()

View File

@ -8,11 +8,11 @@ logger = setup_logger("KnowledgeBase")
class KnowledgeBase: class KnowledgeBase:
def __init__(self, config=None) -> None: def __init__(self, config=None, embed_model=None ) -> None:
self.config = config self.config = config
self._init_config(config) self._init_config(config)
assert embed_model, "embed_model=None"
self.embed_model = EmbeddingModel(config) self.embed_model = embed_model
self.client = MilvusClient("data/vector_base/milvus.db") self.client = MilvusClient("data/vector_base/milvus.db")
def _init_config(self, config): def _init_config(self, config):

View File

@ -1,5 +1,8 @@
from core.startup import dbm, model from core.startup import dbm, model
from models.embedding import ReRanker from models.embedding import ReRanker
from utils.logging_config import setup_logger
logger = setup_logger("server-common")
class Retriever: class Retriever:
@ -30,10 +33,10 @@ class Retriever:
kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res]) kb_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in kb_res])
external += f"知识库信息: \n\n{kb_text}" external += f"知识库信息: \n\n{kb_text}"
# db_res = refs.get("graph_base").get("results", []) db_res = refs.get("graph_base").get("results", [])
# if len(db_res) > 0: if len(db_res) > 0:
# db_text = "\n".join([f"{r['id']}: {r['entity']['text']}" for r in db_res]) 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}" external += f"图数据库信息: \n\n{db_text}"
if len(external) > 0: if len(external) > 0:
query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}" query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}"
@ -53,9 +56,9 @@ class Retriever:
results = [] results = []
if meta.get("use_graph"): if meta.get("use_graph"):
for entitie in entities: for entitie in entities:
result = dbm.graph_base.query_entity_like(entitie) result = dbm.graph_base.query_by_vector(entitie)
results.extend(result) if result else None if result != []:
results.extend(result)
return {"results": self.format_query_results(results)} return {"results": self.format_query_results(results)}
def query_knowledgebase(self, query, history, meta): def query_knowledgebase(self, query, history, meta):
@ -113,29 +116,44 @@ class Retriever:
return rewritten_query, entities return rewritten_query, entities
def format_query_results(self, results): def format_query_results(sfle, results):
formatted_results = {"nodes": [], "edges": []} formatted_results = {"nodes": [], "edges": []}
for row in results:
n, relations, m = row node_dict = {}
formatted_results["nodes"].append({
"id": n.id, for item in results:
"name": n._properties["name"], if isinstance(item[1], list) and len(item[1]) > 0:
"properties": n._properties relationship = item[1][0]
}) rel_id = relationship.element_id
formatted_results["nodes"].append({ nodes = relationship.nodes
"id": m.id, if len(nodes) == 2:
"name": m._properties["name"], node1, node2 = nodes
"properties": m._properties
}) node1_id = node1.element_id
for rel in relations: 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({ formatted_results["edges"].append({
"id": rel.id, "id": rel_id,
"type": rel.type, "type": relationship_type,
"source": rel.start_node.id, "source_id": node1_id,
"target": rel.end_node.id, "target_id": node2_id,
"source_name": rel.start_node._properties["name"], "source_name": node1_name,
"target_name": rel.end_node._properties["name"], "target_name": node2_name
}) })
formatted_results["nodes"] = list(node_dict.values())
return formatted_results return formatted_results
def __call__(self, query, history, meta): def __call__(self, query, history, meta):