diff --git a/src/config/static/models.py b/src/config/static/models.py index 75d6856e..9d4d4b86 100644 --- a/src/config/static/models.py +++ b/src/config/static/models.py @@ -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, ), } diff --git a/src/knowledge/implementations/milvus.py b/src/knowledge/implementations/milvus.py index 862f72b9..4d254f86 100644 --- a/src/knowledge/implementations/milvus.py +++ b/src/knowledge/implementations/milvus.py @@ -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 集合""" diff --git a/src/knowledge/services/upload_graph_service.py b/src/knowledge/services/upload_graph_service.py index da3556ae..b9b69716 100644 --- a/src/knowledge/services/upload_graph_service.py +++ b/src/knowledge/services/upload_graph_service.py @@ -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] diff --git a/src/models/embed.py b/src/models/embed.py index a63bbc58..f05ca5a7 100644 --- a/src/models/embed.py +++ b/src/models/embed.py @@ -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: diff --git a/src/services/evaluation_service.py b/src/services/evaluation_service.py index 1c4b42b9..a1bc18b4 100644 --- a/src/services/evaluation_service.py +++ b/src/services/evaluation_service.py @@ -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):