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:
肖泽涛 2026-02-28 13:09:53 +08:00
parent 3b11331f9a
commit 0695f6f4cd
5 changed files with 41 additions and 13 deletions

View File

@ -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,
),
}

View File

@ -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 集合"""

View File

@ -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]

View File

@ -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:

View File

@ -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):