"""嵌入模型抽象层 — 策略模式. 支持的 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")) # 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 logging import os import threading from typing import Protocol from src.core.config import EmbedConfig logger = logging.getLogger("md-vector-db") _HF_MIRROR = os.environ.get("HF_MIRROR", "https://hf-mirror.com") _HF_ENV_LOCK = threading.Lock() # -- 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 def dimension(self) -> int: """返回嵌入向量的维度.""" ... def embed(self, texts: list[str]) -> list[list[float]]: """对文本列表进行嵌入. Args: texts: 待嵌入的文本列表 Returns: 嵌入向量列表,每个向量为 float 列表 """ ... # -- 本地模型 -- class LocalEmbedder: """sentence-transformers 本地模型嵌入器. 自动检测 GPU.""" def __init__(self, config: EmbedConfig): from sentence_transformers import SentenceTransformer # 自动检测 GPU try: import torch device = "cuda" if torch.cuda.is_available() else "cpu" if device == "cuda": logger.info("检测到 GPU: %s", torch.cuda.get_device_name(0)) except ImportError: device = "cpu" self._config = config try: self._model = SentenceTransformer( config.local_model, local_files_only=True, device=device ) except Exception: logger.info("模型未缓存, 通过镜像下载 %s", config.local_model) # 通过 HF_ENDPOINT 环境变量设置镜像(sentence-transformers 依赖 huggingface_hub) # 临时设置仅用于模型下载,下载完成后还原 old_endpoint = os.environ.get("HF_ENDPOINT") with _HF_ENV_LOCK: os.environ["HF_ENDPOINT"] = _HF_MIRROR try: self._model = SentenceTransformer( config.local_model, device=device ) finally: with _HF_ENV_LOCK: if old_endpoint is not None: os.environ["HF_ENDPOINT"] = old_endpoint else: os.environ.pop("HF_ENDPOINT", None) @property def dimension(self) -> int: return self._model.get_embedding_dimension() def embed(self, texts: list[str]) -> list[list[float]]: if not texts: raise ValueError("文本列表不能为空") embeddings = self._model.encode(texts, normalize_embeddings=True) return embeddings.tolist() # -- API Provider 基类 -- class _BaseAPIEmbedder: """API 嵌入器基类: 统一 api_base/model 解析逻辑.""" 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 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") self._client = None # 懒初始化,首次 embed() 时创建 def _ensure_client(self): """懒创建 OpenAI 客户端(避免 import 时依赖 openai 包).""" if self._client is None: from openai import OpenAI self._client = OpenAI(base_url=self._api_base, api_key=self._api_key) def embed(self, texts: list[str]) -> list[list[float]]: if not texts: raise ValueError("文本列表不能为空") self._ensure_client() response = self._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 # 惰性导入(仅 DashScope 使用) 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": [...]}, ...]}} output = data.get("output") if output is None: raise ValueError(f"DashScope 响应缺少 output 字段: {data}") embeddings_raw = output.get("embeddings") if not isinstance(embeddings_raw, list): raise ValueError(f"DashScope embeddings 不是列表: {type(embeddings_raw)}") # 按 text_index 排序确保顺序 embeddings_raw_sorted = sorted(embeddings_raw, key=lambda x: x.get("text_index", 0)) return [e["embedding"] for e in embeddings_raw_sorted] # -- 工厂函数 -- _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) 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( 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