110 lines
3.4 KiB
Python
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)
|