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")
|
base_url: str = Field(..., description="API 基础 URL")
|
||||||
api_key: str = Field(..., description="API Key 或环境变量名")
|
api_key: str = Field(..., description="API Key 或环境变量名")
|
||||||
model_id: str | None = Field(None, description="可选的模型 ID")
|
model_id: str | None = Field(None, description="可选的模型 ID")
|
||||||
|
batch_size: int = Field(40, description="批量向量化大小")
|
||||||
|
|
||||||
|
|
||||||
class RerankerInfo(BaseModel):
|
class RerankerInfo(BaseModel):
|
||||||
@ -203,6 +204,7 @@ DEFAULT_EMBED_MODELS: dict[str, EmbedModelInfo] = {
|
|||||||
dimension=1024,
|
dimension=1024,
|
||||||
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings",
|
base_url="https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings",
|
||||||
api_key="DASHSCOPE_API_KEY",
|
api_key="DASHSCOPE_API_KEY",
|
||||||
|
batch_size=10,
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -190,13 +190,14 @@ 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)
|
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):
|
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)
|
||||||
|
batch_size = int(getattr(embedding_model, "batch_size", 40) or 40)
|
||||||
return partial(embedding_model.batch_encode, batch_size=40)
|
return partial(embedding_model.batch_encode, batch_size=batch_size)
|
||||||
|
|
||||||
async def _get_milvus_collection(self, db_id: str):
|
async def _get_milvus_collection(self, db_id: str):
|
||||||
"""获取或创建 Milvus 集合"""
|
"""获取或创建 Milvus 集合"""
|
||||||
|
|||||||
@ -488,17 +488,24 @@ class UploadGraphService:
|
|||||||
logger.error(f"加载图数据库信息失败:{e}")
|
logger.error(f"加载图数据库信息失败:{e}")
|
||||||
return False
|
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):
|
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
|
return outputs
|
||||||
else:
|
else:
|
||||||
outputs = await self.embed_model.aencode(text)
|
outputs = await self.embed_model.aencode(text)
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def get_embedding(self, text, batch_size=40):
|
def get_embedding(self, text, batch_size=None):
|
||||||
if isinstance(text, list):
|
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
|
return outputs
|
||||||
else:
|
else:
|
||||||
outputs = self.embed_model.encode([text])[0]
|
outputs = self.embed_model.encode([text])[0]
|
||||||
|
|||||||
@ -11,7 +11,17 @@ from src.utils import get_docker_safe_url, hashstr, logger
|
|||||||
|
|
||||||
|
|
||||||
class BaseEmbeddingModel(ABC):
|
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:
|
Args:
|
||||||
model: 模型名称,冗余设计,同name
|
model: 模型名称,冗余设计,同name
|
||||||
@ -20,12 +30,14 @@ class BaseEmbeddingModel(ABC):
|
|||||||
url: 请求URL,冗余设计,同base_url
|
url: 请求URL,冗余设计,同base_url
|
||||||
base_url: 基础URL,请求URL,冗余设计,同url
|
base_url: 基础URL,请求URL,冗余设计,同url
|
||||||
api_key: 请求API密钥
|
api_key: 请求API密钥
|
||||||
|
batch_size: 模型推荐的批量向量化大小
|
||||||
"""
|
"""
|
||||||
base_url = base_url or url
|
base_url = base_url or url
|
||||||
self.model = model or name
|
self.model = model or name
|
||||||
self.dimension = dimension
|
self.dimension = dimension
|
||||||
self.base_url = get_docker_safe_url(base_url)
|
self.base_url = get_docker_safe_url(base_url)
|
||||||
self.api_key = os.getenv(api_key, api_key)
|
self.api_key = os.getenv(api_key, api_key)
|
||||||
|
self.batch_size = int(batch_size or 40)
|
||||||
self.embed_state = {}
|
self.embed_state = {}
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
@ -46,8 +58,9 @@ 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 | None = None) -> list[list[float]]:
|
||||||
# logger.info(f"Batch encoding {len(messages)} messages")
|
# logger.info(f"Batch encoding {len(messages)} messages")
|
||||||
|
batch_size = batch_size or self.batch_size
|
||||||
data = []
|
data = []
|
||||||
task_id = None
|
task_id = None
|
||||||
if len(messages) > batch_size:
|
if len(messages) > batch_size:
|
||||||
@ -67,7 +80,8 @@ 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 | None = None) -> list[list[float]]:
|
||||||
|
batch_size = batch_size or self.batch_size
|
||||||
data = []
|
data = []
|
||||||
task_id = None
|
task_id = None
|
||||||
if len(messages) > batch_size:
|
if len(messages) > batch_size:
|
||||||
|
|||||||
@ -5,6 +5,7 @@ import uuid
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
from src import config
|
||||||
from src.knowledge import knowledge_base
|
from src.knowledge import knowledge_base
|
||||||
from src.models import select_model
|
from src.models import select_model
|
||||||
from src.repositories.evaluation_repository import EvaluationRepository
|
from src.repositories.evaluation_repository import EvaluationRepository
|
||||||
@ -345,9 +346,9 @@ class EvaluationService:
|
|||||||
|
|
||||||
await context.set_progress(15, "向量化")
|
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:
|
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 ""
|
embedding_model_id = embed_info.get("name") or embed_info.get("model") or ""
|
||||||
if not embedding_model_id:
|
if not embedding_model_id:
|
||||||
raise ValueError("Embedding model not specified")
|
raise ValueError("Embedding model not specified")
|
||||||
@ -355,11 +356,14 @@ class EvaluationService:
|
|||||||
from src.models import select_embedding_model, select_model
|
from src.models import select_embedding_model, select_model
|
||||||
|
|
||||||
embed_model = select_embedding_model(embedding_model_id)
|
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
|
# TODO: Performance Optimization
|
||||||
# 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=batch_size)
|
||||||
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