Fix embedding 400 errors
This commit is contained in:
parent
4d5f38beae
commit
9cd6470849
@ -191,13 +191,13 @@ class MilvusKB(KnowledgeBase):
|
|||||||
def _get_async_embedding_function(self, embed_info: dict):
|
def _get_async_embedding_function(self, embed_info: dict):
|
||||||
"""获取 embedding 函数"""
|
"""获取 embedding 函数"""
|
||||||
embedding_model = self._get_async_embedding(embed_info)
|
embedding_model = self._get_async_embedding(embed_info)
|
||||||
return partial(embedding_model.abatch_encode, batch_size=40)
|
return partial(embedding_model.abatch_encode, batch_size=10)
|
||||||
|
|
||||||
def _get_embedding_function(self, embed_info: dict):
|
def _get_embedding_function(self, embed_info: dict):
|
||||||
"""获取 embedding 函数"""
|
"""获取 embedding 函数"""
|
||||||
embedding_model = self._get_async_embedding(embed_info)
|
embedding_model = self._get_async_embedding(embed_info)
|
||||||
|
|
||||||
return partial(embedding_model.batch_encode, batch_size=40)
|
return partial(embedding_model.batch_encode, batch_size=10)
|
||||||
|
|
||||||
async def _get_milvus_collection(self, db_id: str):
|
async def _get_milvus_collection(self, db_id: str):
|
||||||
"""获取或创建 Milvus 集合"""
|
"""获取或创建 Milvus 集合"""
|
||||||
|
|||||||
@ -488,7 +488,7 @@ class UploadGraphService:
|
|||||||
logger.error(f"加载图数据库信息失败:{e}")
|
logger.error(f"加载图数据库信息失败:{e}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
async def aget_embedding(self, text, batch_size=40):
|
async def aget_embedding(self, text, batch_size=10):
|
||||||
if isinstance(text, list):
|
if isinstance(text, list):
|
||||||
outputs = await self.embed_model.abatch_encode(text, batch_size=batch_size)
|
outputs = await self.embed_model.abatch_encode(text, batch_size=batch_size)
|
||||||
return outputs
|
return outputs
|
||||||
|
|||||||
@ -46,7 +46,7 @@ class BaseEmbeddingModel(ABC):
|
|||||||
"""等同于aencode"""
|
"""等同于aencode"""
|
||||||
return await self.aencode(queries)
|
return await self.aencode(queries)
|
||||||
|
|
||||||
def batch_encode(self, messages: list[str], batch_size: int = 40) -> list[list[float]]:
|
def batch_encode(self, messages: list[str], batch_size: int = 10) -> list[list[float]]:
|
||||||
# logger.info(f"Batch encoding {len(messages)} messages")
|
# logger.info(f"Batch encoding {len(messages)} messages")
|
||||||
data = []
|
data = []
|
||||||
task_id = None
|
task_id = None
|
||||||
@ -67,24 +67,39 @@ class BaseEmbeddingModel(ABC):
|
|||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
async def abatch_encode(self, messages: list[str], batch_size: int = 40) -> list[list[float]]:
|
async def abatch_encode(self, messages: list[str], batch_size: int = 10) -> list[list[float]]:
|
||||||
data = []
|
data = []
|
||||||
task_id = None
|
task_id = None
|
||||||
if len(messages) > batch_size:
|
if len(messages) > batch_size:
|
||||||
task_id = hashstr(messages)
|
task_id = hashstr(messages)
|
||||||
self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0}
|
self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0}
|
||||||
|
|
||||||
tasks = []
|
#保留原有逻辑:
|
||||||
|
#使用 asyncio.gather 并发执行所有 embedding 批次请求:
|
||||||
|
# tasks = []
|
||||||
|
# for i in range(0, len(messages), batch_size):
|
||||||
|
# group_msg = messages[i : i + batch_size]
|
||||||
|
# tasks.append(self.aencode(group_msg))
|
||||||
|
|
||||||
|
# results = await asyncio.gather(*tasks)
|
||||||
|
# for res in results:
|
||||||
|
# data.extend(res)
|
||||||
|
|
||||||
|
# if task_id:
|
||||||
|
# self.embed_state[task_id]["progress"] = len(messages)
|
||||||
|
# self.embed_state[task_id]["status"] = "completed"
|
||||||
|
|
||||||
|
# return data
|
||||||
|
|
||||||
for i in range(0, len(messages), batch_size):
|
for i in range(0, len(messages), batch_size):
|
||||||
group_msg = messages[i : i + batch_size]
|
group_msg = messages[i : i + batch_size]
|
||||||
tasks.append(self.aencode(group_msg))
|
logger.info(f"Async encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
|
||||||
|
res = await self.aencode(group_msg)
|
||||||
results = await asyncio.gather(*tasks)
|
|
||||||
for res in results:
|
|
||||||
data.extend(res)
|
data.extend(res)
|
||||||
|
if task_id:
|
||||||
|
self.embed_state[task_id]["progress"] = i + len(group_msg)
|
||||||
|
|
||||||
if task_id:
|
if task_id:
|
||||||
self.embed_state[task_id]["progress"] = len(messages)
|
|
||||||
self.embed_state[task_id]["status"] = "completed"
|
self.embed_state[task_id]["status"] = "completed"
|
||||||
|
|
||||||
return data
|
return data
|
||||||
|
|||||||
@ -336,7 +336,7 @@ class EvaluationService:
|
|||||||
# Currently, we re-calculate embeddings for ALL chunks in the KB for every benchmark generation.
|
# Currently, we re-calculate embeddings for ALL chunks in the KB for every benchmark generation.
|
||||||
# This is inefficient for large KBs (O(N) embedding calls).
|
# This is inefficient for large KBs (O(N) embedding calls).
|
||||||
# Optimization: Reuse existing embeddings from Vector DB if embedding_model_id matches the KB's embedding model.
|
# Optimization: Reuse existing embeddings from Vector DB if embedding_model_id matches the KB's embedding model.
|
||||||
embeddings = await embed_model.abatch_encode(contents, batch_size=40)
|
embeddings = await embed_model.abatch_encode(contents, batch_size=10)
|
||||||
norms = [math.sqrt(sum(x * x for x in vec)) or 1.0 for vec in embeddings]
|
norms = [math.sqrt(sum(x * x for x in vec)) or 1.0 for vec in embeddings]
|
||||||
|
|
||||||
def cosine(a, b, na, nb):
|
def cosine(a, b, na, nb):
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user