feat: multi-provider embedder — support OpenAI/DashScope with provider-specific defaults

This commit is contained in:
2026-07-05 01:43:23 +08:00
parent 5672472030
commit b89c6076e1
5 changed files with 173 additions and 36 deletions
+100 -16
View File
@@ -1,12 +1,20 @@
"""嵌入模型抽象层 — 策略模式.
支持的 Provider:
openai — OpenAI / 硅基流动 / 智谱 / DeepSeek / 月之暗面 等 OpenAI 兼容服务
dashscope — 阿里云 DashScope (通义千问)
使用方式:
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:
# OpenAI 兼容 API
embedder = create_embedder(EmbedConfig(mode="api", provider="openai"))
# 阿里云 DashScope
embedder = create_embedder(EmbedConfig(mode="api", provider="dashscope"))
vectors = batch_embed(embedder, long_text_list)
"""
import os
@@ -17,10 +25,24 @@ from src.core.config import EmbedConfig
logger = logging.getLogger(__name__)
# 模型下载镜像 (仅首次下载时使用)
_HF_MIRROR = "https://hf-mirror.com"
# -- Provider 默认配置 --
_PROVIDER_DEFAULTS: dict[str, dict[str, str | int]] = {
"openai": {
"api_base": "https://api.openai.com/v1",
"model": "text-embedding-3-small",
"dimension": 1536,
},
"dashscope": {
"api_base": "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding",
"model": "text-embedding-v4",
"dimension": 1536,
},
}
# -- 接口 --
class Embedder(Protocol):
"""嵌入器接口."""
@property
@@ -28,8 +50,9 @@ class Embedder(Protocol):
def embed(self, texts: list[str]) -> list[list[float]]: ...
# -- 本地模型 --
class LocalEmbedder:
"""本地 sentence-transformers 模型嵌入器."""
"""sentence-transformers 本地模型嵌入器."""
def __init__(self, config: EmbedConfig):
from sentence_transformers import SentenceTransformer
@@ -61,36 +84,97 @@ class LocalEmbedder:
return embeddings.tolist()
class APIEmbedder:
"""OpenAI 兼容 API 嵌入器."""
# -- API Provider 基类 --
class _BaseAPIEmbedder:
"""API 嵌入器基类: 统一 api_base/model 解析逻辑."""
def __init__(self, config: EmbedConfig):
self._api_base = config.api_base
def __init__(self, config: EmbedConfig, provider: str):
defaults = _PROVIDER_DEFAULTS.get(provider, {})
self._api_base = config.api_base or str(defaults.get("api_base", ""))
self._model = config.model or str(defaults.get("model", ""))
self._api_key = config.api_key
self._dimension = int(defaults.get("dimension", 1536))
@property
def dimension(self) -> int:
return 1536
return self._dimension
def embed(self, texts: list[str]) -> list[list[float]]:
raise NotImplementedError
# -- OpenAI 兼容 API --
class OpenAIEmbedder(_BaseAPIEmbedder):
"""OpenAI / 硅基流动 / 智谱 / DeepSeek / 月之暗面 等."""
def __init__(self, config: EmbedConfig):
super().__init__(config, "openai")
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,
)
response = client.embeddings.create(model=self._model, input=texts)
return [d.embedding for d in response.data]
# -- 阿里云 DashScope --
class DashscopeEmbedder(_BaseAPIEmbedder):
"""阿里云 DashScope 嵌入 (自定义 HTTP API, 非 OpenAI 兼容)."""
def __init__(self, config: EmbedConfig):
super().__init__(config, "dashscope")
def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
raise ValueError("文本列表不能为空")
import requests
resp = requests.post(
self._api_base,
headers={
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
},
json={
"model": self._model,
"input": {"texts": texts},
},
timeout=60,
)
resp.raise_for_status()
data = resp.json()
# DashScope 返回: {"output": {"embeddings": [{"text_index": 0, "embedding": [...]}, ...]}}
embeddings_raw = data.get("output", {}).get("embeddings", [])
# 按 text_index 排序确保顺序
embeddings_raw.sort(key=lambda x: x.get("text_index", 0))
return [e["embedding"] for e in embeddings_raw]
# -- 工厂函数 --
_PROVIDER_CLASSES: dict[str, type[_BaseAPIEmbedder]] = {
"openai": OpenAIEmbedder,
"dashscope": DashscopeEmbedder,
}
SUPPORTED_PROVIDERS = list(_PROVIDER_CLASSES.keys())
def create_embedder(config: EmbedConfig) -> Embedder:
"""工厂函数: 根据配置创建嵌入器."""
if config.mode == "local":
return LocalEmbedder(config)
elif config.mode == "api":
return APIEmbedder(config)
else:
raise ValueError(f"不支持的嵌入模式: {config.mode}")
if config.mode == "api":
provider = config.provider or "openai"
cls = _PROVIDER_CLASSES.get(provider)
if cls is None:
raise ValueError(
f"不支持的 provider: {provider}, 可选: {SUPPORTED_PROVIDERS}"
)
return cls(config)
raise ValueError(f"不支持的嵌入模式: {config.mode}")
def batch_embed(