feat: multi-provider embedder — support OpenAI/DashScope with provider-specific defaults
This commit is contained in:
+100
-16
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user