Initial commit
This commit is contained in:
@@ -0,0 +1,109 @@
|
||||
"""本地 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)
|
||||
Reference in New Issue
Block a user