refactor: strategy-pattern embedder (Local/API), remove global env vars, add batch embedding

This commit is contained in:
2026-07-05 01:35:54 +08:00
parent 81a8173a38
commit ad9e5ce449
3 changed files with 112 additions and 66 deletions
+77 -52
View File
@@ -1,79 +1,104 @@
"""嵌入模型抽象层."""
"""嵌入模型抽象层 — 策略模式.
使用方式:
from src.core.config import EmbedConfig
from src.core.embedder import create_embedder, batch_embed
embedder = create_embedder(EmbedConfig(mode="local"))
vectors = embedder.embed(["文本1", "文本2"])
# 或分批嵌入避免 OOM:
vectors = batch_embed(embedder, long_text_list)
"""
import os
# 解决 Windows SSL 证书问题: 强制使用本地缓存模型, 避免联网验证
# 模型首次下载需要先临时取消这些环境变量 (或手动运行一次下载脚本)
os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
from sentence_transformers import SentenceTransformer
import logging
from typing import Protocol
from src.core.config import EmbedConfig
logger = logging.getLogger(__name__)
class Embedder:
"""文本嵌入器, 支持本地模型和 API 两种模式."""
# 模型下载镜像 (仅首次下载时使用)
_HF_MIRROR = "https://hf-mirror.com"
class Embedder(Protocol):
"""嵌入器接口."""
@property
def dimension(self) -> int: ...
def embed(self, texts: list[str]) -> list[list[float]]: ...
class LocalEmbedder:
"""本地 sentence-transformers 模型嵌入器."""
def __init__(self, config: EmbedConfig):
from sentence_transformers import SentenceTransformer
self._config = config
if config.mode == "local":
# 优先从本地缓存加载, 若模型未下载则通过镜像下载
try:
self._model = SentenceTransformer(
config.local_model, local_files_only=True
)
except Exception:
logger.info("模型未缓存, 通过镜像下载 %s", config.local_model)
old_endpoint = os.environ.get("HF_ENDPOINT")
os.environ["HF_ENDPOINT"] = _HF_MIRROR
try:
self._model = SentenceTransformer(
config.local_model, local_files_only=True
)
except Exception:
# 模型未缓存: 临时开启网络, 使用国内镜像
os.environ.pop("HF_HUB_OFFLINE", None)
os.environ.pop("TRANSFORMERS_OFFLINE", None)
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
self._model = SentenceTransformer(config.local_model)
# 恢复离线设置, 后续加载走缓存
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
elif config.mode == "api":
self._model = None # 延迟初始化, 需要 openai 包
self._api_base = config.api_base
self._api_key = config.api_key
else:
raise ValueError(f"不支持的嵌入模式: {config.mode}")
@property
def mode(self) -> str:
"""当前嵌入模式."""
return self._config.mode
finally:
if old_endpoint is not None:
os.environ["HF_ENDPOINT"] = old_endpoint
else:
os.environ.pop("HF_ENDPOINT", None)
@property
def dimension(self) -> int:
"""嵌入向量维度."""
if self.mode == "local":
return self._model.get_embedding_dimension()
else:
# 默认 OpenAI text-embedding-ada-002 / text-embedding-3-small 维度
return 1536
return self._model.get_embedding_dimension()
def embed(self, texts: list[str]) -> list[list[float]]:
"""对文本列表进行嵌入, 返回向量列表."""
if not texts:
raise ValueError("文本列表不能为空")
if self.mode == "local":
embeddings = self._model.encode(texts, normalize_embeddings=True)
return embeddings.tolist()
else:
return self._embed_via_api(texts)
embeddings = self._model.encode(texts, normalize_embeddings=True)
return embeddings.tolist()
def _embed_via_api(self, texts: list[str]) -> list[list[float]]:
"""通过 OpenAI 兼容 API 嵌入(延迟导入 openai."""
class APIEmbedder:
"""OpenAI 兼容 API 嵌入器."""
def __init__(self, config: EmbedConfig):
self._api_base = config.api_base
self._api_key = config.api_key
@property
def dimension(self) -> int:
return 1536
def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
raise ValueError("文本列表不能为空")
from openai import OpenAI
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
response = client.embeddings.create(
model="text-embedding-3-small",
input=texts,
model="text-embedding-3-small", input=texts,
)
return [d.embedding for d in response.data]
def create_embedder(config: EmbedConfig) -> Embedder:
"""工厂函数: 根据配置创建嵌入器."""
return Embedder(config)
if config.mode == "local":
return LocalEmbedder(config)
elif config.mode == "api":
return APIEmbedder(config)
else:
raise ValueError(f"不支持的嵌入模式: {config.mode}")
def batch_embed(
embedder: Embedder, texts: list[str], batch_size: int = 32
) -> list[list[float]]:
"""分批嵌入, 避免一次性传入过多文本导致 OOM."""
all_embeddings = []
for i in range(0, len(texts), batch_size):
batch = texts[i:i + batch_size]
all_embeddings.extend(embedder.embed(batch))
return all_embeddings