优化embedding阶段的并行处理
This commit is contained in:
parent
b66fbf9b38
commit
9c61a8a436
@ -21,7 +21,7 @@ async def chat_get():
|
||||
return "Chat Get!"
|
||||
|
||||
@chat.post("/")
|
||||
def chat_post(
|
||||
async def chat_post(
|
||||
query: str = Body(...),
|
||||
meta: dict = Body(None),
|
||||
history: list[dict] | None = Body(None),
|
||||
|
||||
@ -176,6 +176,6 @@ async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str]
|
||||
if not file_path.endswith('.jsonl'):
|
||||
raise HTTPException(status_code=400, detail="file_path must be a jsonl file")
|
||||
|
||||
graph_base.jsonl_file_add_entity(file_path, kgdb_name)
|
||||
await graph_base.jsonl_file_add_entity(file_path, kgdb_name)
|
||||
return {"message": "Entity successfully added"}
|
||||
|
||||
|
||||
@ -106,7 +106,7 @@ class GraphDatabase:
|
||||
with self.driver.session() as session:
|
||||
session.execute_write(create, triples)
|
||||
|
||||
def txt_add_vector_entity(self, triples, kgdb_name='neo4j'):
|
||||
async def txt_add_vector_entity(self, triples, kgdb_name='neo4j'):
|
||||
"""添加实体三元组"""
|
||||
self.use_database(kgdb_name)
|
||||
def _index_exists(tx, index_name):
|
||||
@ -140,6 +140,14 @@ class GraphDatabase:
|
||||
}} }};
|
||||
""")
|
||||
|
||||
def _batch_set_embeddings(tx, entity_embedding_pairs):
|
||||
"""批量设置实体的嵌入向量"""
|
||||
for entity_name, embedding in entity_embedding_pairs:
|
||||
tx.run("""
|
||||
MATCH (e:Entity {name: $name})
|
||||
CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding)
|
||||
""", name=entity_name, embedding=embedding)
|
||||
|
||||
# 判断模型名称是否匹配
|
||||
cur_embed_info = config.embed_model_names[config.embed_model]
|
||||
self.embed_model_name = self.embed_model_name or cur_embed_info.get('name')
|
||||
@ -151,18 +159,36 @@ class GraphDatabase:
|
||||
session.execute_write(_create_graph, triples)
|
||||
logger.info(f"Creating vector index for {kgdb_name} with {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)
|
||||
|
||||
# 收集所有需要处理的实体名称,去重
|
||||
all_entities = []
|
||||
for entry in triples:
|
||||
if entry['h'] not in all_entities:
|
||||
all_entities.append(entry['h'])
|
||||
if entry['t'] not in all_entities:
|
||||
all_entities.append(entry['t'])
|
||||
|
||||
# 批量处理实体
|
||||
max_batch_size = 1024 # 限制此部分的主要是内存大小 1024 * 1024 * 4 / 1024 / 1024 = 4GB
|
||||
total_entities = len(all_entities)
|
||||
|
||||
for i in range(0, total_entities, max_batch_size):
|
||||
batch_entities = all_entities[i:i+max_batch_size]
|
||||
logger.debug(f"Processing entities batch {i//max_batch_size + 1}/{(total_entities-1)//max_batch_size + 1} ({len(batch_entities)} entities)")
|
||||
|
||||
# 批量获取嵌入向量
|
||||
batch_embeddings = await self.aget_embedding(batch_entities)
|
||||
|
||||
# 将实体名称和嵌入向量配对
|
||||
entity_embedding_pairs = list(zip(batch_entities, batch_embeddings))
|
||||
|
||||
# 批量写入数据库
|
||||
session.execute_write(_batch_set_embeddings, entity_embedding_pairs)
|
||||
|
||||
# 数据添加完成后保存图信息
|
||||
self.save_graph_info()
|
||||
|
||||
def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'):
|
||||
async def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'):
|
||||
self.status = "processing"
|
||||
kgdb_name = kgdb_name or 'neo4j'
|
||||
self.use_database(kgdb_name) # 切换到指定数据库
|
||||
@ -176,7 +202,7 @@ class GraphDatabase:
|
||||
|
||||
triples = list(read_triples(file_path))
|
||||
|
||||
self.txt_add_vector_entity(triples, kgdb_name)
|
||||
await self.txt_add_vector_entity(triples, kgdb_name)
|
||||
|
||||
self.status = "open"
|
||||
# 更新并保存图数据库信息
|
||||
@ -355,9 +381,23 @@ class GraphDatabase:
|
||||
with self.driver.session() as session:
|
||||
return session.execute_read(query, node_name, hops)
|
||||
|
||||
async def aget_embedding(self, text):
|
||||
from src import knowledge_base
|
||||
|
||||
if isinstance(text, list):
|
||||
outputs = await knowledge_base.embed_model.abatch_encode(text, batch_size=40)
|
||||
return outputs
|
||||
else:
|
||||
outputs = await knowledge_base.embed_model.aencode(text)
|
||||
return outputs
|
||||
|
||||
def get_embedding(self, text):
|
||||
with torch.no_grad():
|
||||
from src import knowledge_base
|
||||
from src import knowledge_base
|
||||
|
||||
if isinstance(text, list):
|
||||
outputs = knowledge_base.embed_model.batch_encode(text, batch_size=40)
|
||||
return outputs
|
||||
else:
|
||||
outputs = knowledge_base.embed_model.encode([text])[0]
|
||||
return outputs
|
||||
|
||||
|
||||
@ -1,6 +1,7 @@
|
||||
import os
|
||||
import json
|
||||
import requests
|
||||
import asyncio
|
||||
from FlagEmbedding import FlagModel
|
||||
from zhipuai import ZhipuAI
|
||||
|
||||
@ -26,6 +27,15 @@ class BaseEmbeddingModel:
|
||||
def encode_queries(self, queries):
|
||||
return self.predict(queries)
|
||||
|
||||
async def aencode(self, message):
|
||||
return await asyncio.to_thread(self.encode, message)
|
||||
|
||||
async def aencode_queries(self, queries):
|
||||
return await asyncio.to_thread(self.encode_queries, queries)
|
||||
|
||||
async def abatch_encode(self, messages, batch_size=20):
|
||||
return await asyncio.to_thread(self.batch_encode, messages, batch_size)
|
||||
|
||||
def batch_encode(self, messages, batch_size=20):
|
||||
logger.info(f"Batch encoding {len(messages)} messages")
|
||||
data = []
|
||||
|
||||
Loading…
Reference in New Issue
Block a user