"""本地 BGE-M3 嵌入实现(可选)。 需要额外安装本地依赖: uv sync --extra local-embedding 模型 BAAI/bge-m3 首次加载约 2.3 GB(从 HuggingFace 自动下载)。 默认 1024 维,与 DashScope 兼容。 """ from __future__ import annotations import asyncio import os from typing import TYPE_CHECKING from loguru import logger from .base import AsyncEmbeddingProvider, EmbeddingProvider from .models import EmbeddingError if TYPE_CHECKING: from sentence_transformers import SentenceTransformer LOCAL_DEFAULT_MODEL = "BAAI/bge-m3" LOCAL_DEFAULT_DIM = 1024 def _read_env(key: str, default: str | None = None) -> str | None: val = os.environ.get(key) if val is None or val.strip() == "": return default return val.strip() def _try_import_st() -> type[SentenceTransformer]: try: from sentence_transformers import SentenceTransformer # noqa: F811 except ImportError as e: raise EmbeddingError( "本地 BGE-M3 后端需要 sentence-transformers,请执行: " "uv sync --extra local-embedding" ) from e return SentenceTransformer class LocalBGEEmbeddingProvider(EmbeddingProvider): """本地 BGE-M3 同步实现。""" name = "local-bge" def __init__( self, *, model: str | None = None, device: str | None = None, normalize: bool = True, ) -> None: st_cls = _try_import_st() # 模型优先级:CLI/参数 > LOCAL_EMBEDDING_MODEL > 默认值 # 不读全局 EMBEDDING_MODEL,避免与 DashScope 冲突 self.model = ( model or _read_env("LOCAL_EMBEDDING_MODEL", LOCAL_DEFAULT_MODEL) or LOCAL_DEFAULT_MODEL ) self.dim = LOCAL_DEFAULT_DIM self._device = device # None -> 让 ST 自选 cpu/cuda self._normalize = normalize logger.info("加载本地嵌入模型 {} (device={}),首次会下载...", self.model, device or "auto") self._st = st_cls(self.model, device=device) # 实际维度自检 actual_dim = self._st.get_sentence_embedding_dimension() if actual_dim and actual_dim != self.dim: logger.warning( "BGE 模型实际维度 {} 与默认 {} 不一致,以实际为准", actual_dim, self.dim ) self.dim = actual_dim def embed_batch(self, texts: list[str]) -> list[list[float]]: if not texts: return [] try: arr = self._st.encode( texts, normalize_embeddings=self._normalize, convert_to_numpy=True, show_progress_bar=False, ) except Exception as e: # noqa: BLE001 raise EmbeddingError(f"BGE-M3 推理失败: {type(e).__name__}: {e}") from e return [v.tolist() for v in arr] class LocalBGEAsyncEmbeddingProvider(AsyncEmbeddingProvider): """本地 BGE-M3 异步包装。 sentence-transformers 本身是同步的,这里用 asyncio.to_thread 包一层, 主要为了让批处理脚本能用统一的 async 接口。 """ name = "local-bge" def __init__(self, **kwargs: object) -> None: self._sync = LocalBGEEmbeddingProvider(**kwargs) # type: ignore[arg-type] self.model = self._sync.model self.dim = self._sync.dim async def embed_batch(self, texts: list[str]) -> list[list[float]]: return await asyncio.to_thread(self._sync.embed_batch, texts)