feat: add embedder with local/API dual mode

This commit is contained in:
2026-07-05 01:00:31 +08:00
parent b756c53e25
commit f9e5c37d3e
2 changed files with 110 additions and 0 deletions
+59
View File
@@ -0,0 +1,59 @@
"""嵌入模型抽象层."""
from sentence_transformers import SentenceTransformer
from src.core.config import EmbedConfig
class Embedder:
"""文本嵌入器, 支持本地模型和 API 两种模式."""
def __init__(self, config: EmbedConfig):
self._config = config
if config.mode == "local":
self._model = SentenceTransformer(config.local_model)
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
@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
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)
def _embed_via_api(self, texts: list[str]) -> list[list[float]]:
"""通过 OpenAI 兼容 API 嵌入(延迟导入 openai."""
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,
)
return [d.embedding for d in response.data]
def create_embedder(config: EmbedConfig) -> Embedder:
"""工厂函数: 根据配置创建嵌入器."""
return Embedder(config)