fix(embedding): use model-level batch_size for provider-specific limits
Root cause: DashScope embeddings accepts at most 10 inputs per request, but code used fixed batch_size=40, causing 400 Bad Request during indexing. Changes: - add batch_size to EmbedModelInfo with default 40 - set dashscope/text-embedding-v4 batch_size to 10 - persist batch_size in BaseEmbeddingModel and use it as default in batch/abatch encode - update Milvus, evaluation, and upload-graph embedding paths to read model-configured batch_size
This commit is contained in:
parent
3b11331f9a
commit
0695f6f4cd
@ -30,6 +30,7 @@ class EmbedModelInfo(BaseModel):
|
||||
base_url: str = Field(..., description="API 基础 URL")
|
||||
api_key: str = Field(..., description="API Key 或环境变量名")
|
||||
model_id: str | None = Field(None, description="可选的模型 ID")
|
||||
batch_size: int = Field(40, description="批量向量化大小")
|
||||
|
||||
|
||||
class RerankerInfo(BaseModel):
|
||||
@ -203,6 +204,7 @@ DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
|
||||
dimension=1024,
|
||||
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings",
|
||||
api_key="DASHSCOPE_API_KEY",
|
||||
batch_size=10,
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@ -190,13 +190,14 @@ class MilvusKB(KnowledgeBase):
|
||||
def _get_async_embedding_function(self, embed_info: dict):
|
||||
"""获取 embedding 函数"""
|
||||
embedding_model = self._get_async_embedding(embed_info)
|
||||
return partial(embedding_model.abatch_encode, batch_size=40)
|
||||
batch_size = int(getattr(embedding_model, "batch_size", 40) or 40)
|
||||
return partial(embedding_model.abatch_encode, batch_size=batch_size)
|
||||
|
||||
def _get_embedding_function(self, embed_info: dict):
|
||||
"""获取 embedding 函数"""
|
||||
embedding_model = self._get_async_embedding(embed_info)
|
||||
|
||||
return partial(embedding_model.batch_encode, batch_size=40)
|
||||
batch_size = int(getattr(embedding_model, "batch_size", 40) or 40)
|
||||
return partial(embedding_model.batch_encode, batch_size=batch_size)
|
||||
|
||||
async def _get_milvus_collection(self, db_id: str):
|
||||
"""获取或创建 Milvus 集合"""
|
||||
|
||||
@ -488,17 +488,24 @@ class UploadGraphService:
|
||||
logger.error(f"加载图数据库信息失败:{e}")
|
||||
return False
|
||||
|
||||
async def aget_embedding(self, text, batch_size=40):
|
||||
def _resolve_embedding_batch_size(self, batch_size=None):
|
||||
if batch_size is not None:
|
||||
return batch_size
|
||||
return int(getattr(self.embed_model, "batch_size", 40) or 40)
|
||||
|
||||
async def aget_embedding(self, text, batch_size=None):
|
||||
if isinstance(text, list):
|
||||
outputs = await self.embed_model.abatch_encode(text, batch_size=batch_size)
|
||||
resolved_batch_size = self._resolve_embedding_batch_size(batch_size)
|
||||
outputs = await self.embed_model.abatch_encode(text, batch_size=resolved_batch_size)
|
||||
return outputs
|
||||
else:
|
||||
outputs = await self.embed_model.aencode(text)
|
||||
return outputs
|
||||
|
||||
def get_embedding(self, text, batch_size=40):
|
||||
def get_embedding(self, text, batch_size=None):
|
||||
if isinstance(text, list):
|
||||
outputs = self.embed_model.batch_encode(text, batch_size=batch_size)
|
||||
resolved_batch_size = self._resolve_embedding_batch_size(batch_size)
|
||||
outputs = self.embed_model.batch_encode(text, batch_size=resolved_batch_size)
|
||||
return outputs
|
||||
else:
|
||||
outputs = self.embed_model.encode([text])[0]
|
||||
|
||||
@ -11,7 +11,17 @@ from src.utils import get_docker_safe_url, hashstr, logger
|
||||
|
||||
|
||||
class BaseEmbeddingModel(ABC):
|
||||
def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None, model_id=None):
|
||||
def __init__(
|
||||
self,
|
||||
model=None,
|
||||
name=None,
|
||||
dimension=None,
|
||||
url=None,
|
||||
base_url=None,
|
||||
api_key=None,
|
||||
model_id=None,
|
||||
batch_size=40,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
model: 模型名称,冗余设计,同name
|
||||
@ -20,12 +30,14 @@ class BaseEmbeddingModel(ABC):
|
||||
url: 请求URL,冗余设计,同base_url
|
||||
base_url: 基础URL,请求URL,冗余设计,同url
|
||||
api_key: 请求API密钥
|
||||
batch_size: 模型推荐的批量向量化大小
|
||||
"""
|
||||
base_url = base_url or url
|
||||
self.model = model or name
|
||||
self.dimension = dimension
|
||||
self.base_url = get_docker_safe_url(base_url)
|
||||
self.api_key = os.getenv(api_key, api_key)
|
||||
self.batch_size = int(batch_size or 40)
|
||||
self.embed_state = {}
|
||||
|
||||
@abstractmethod
|
||||
@ -46,8 +58,9 @@ class BaseEmbeddingModel(ABC):
|
||||
"""等同于aencode"""
|
||||
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 | None = None) -> list[list[float]]:
|
||||
# logger.info(f"Batch encoding {len(messages)} messages")
|
||||
batch_size = batch_size or self.batch_size
|
||||
data = []
|
||||
task_id = None
|
||||
if len(messages) > batch_size:
|
||||
@ -67,7 +80,8 @@ class BaseEmbeddingModel(ABC):
|
||||
|
||||
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 | None = None) -> list[list[float]]:
|
||||
batch_size = batch_size or self.batch_size
|
||||
data = []
|
||||
task_id = None
|
||||
if len(messages) > batch_size:
|
||||
|
||||
@ -5,6 +5,7 @@ import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from src import config
|
||||
from src.knowledge import knowledge_base
|
||||
from src.models import select_model
|
||||
from src.repositories.evaluation_repository import EvaluationRepository
|
||||
@ -345,9 +346,9 @@ class EvaluationService:
|
||||
|
||||
await context.set_progress(15, "向量化")
|
||||
|
||||
db_meta = kb_instance.databases_meta.get(db_id, {})
|
||||
embed_info = db_meta.get("embed_info", {})
|
||||
if not embedding_model_id:
|
||||
db_meta = kb_instance.databases_meta.get(db_id, {})
|
||||
embed_info = db_meta.get("embed_info", {})
|
||||
embedding_model_id = embed_info.get("name") or embed_info.get("model") or ""
|
||||
if not embedding_model_id:
|
||||
raise ValueError("Embedding model not specified")
|
||||
@ -355,11 +356,14 @@ class EvaluationService:
|
||||
from src.models import select_embedding_model, select_model
|
||||
|
||||
embed_model = select_embedding_model(embedding_model_id)
|
||||
batch_size = int(getattr(embed_model, "batch_size", 40) or 40)
|
||||
if embedding_model_id in config.embed_model_names:
|
||||
batch_size = config.embed_model_names[embedding_model_id].batch_size
|
||||
# TODO: Performance Optimization
|
||||
# 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).
|
||||
# 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=batch_size)
|
||||
norms = [math.sqrt(sum(x * x for x in vec)) or 1.0 for vec in embeddings]
|
||||
|
||||
def cosine(a, b, na, nb):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user