Files
2026-07-18 15:51:01 +08:00

110 lines
3.4 KiB
Python

"""本地 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)