Merge branch 'dev' of https://github.com/xerrors/ProjectAthena into dev
This commit is contained in:
commit
590dfd1836
3
.gitignore
vendored
3
.gitignore
vendored
@ -29,4 +29,5 @@ cache
|
||||
|
||||
*.pdf
|
||||
src/data
|
||||
neo4j*
|
||||
neo4j*
|
||||
*/package-lock.json
|
||||
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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):
|
||||
@ -117,29 +120,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):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user