Files

267 lines
9.0 KiB
Python
Raw Permalink Normal View History

2026-04-10 16:52:07 +08:00
"""
Embedding 封装:通过 OpenAI 兼容 ``/v1/embeddings`` 远程 API 获取向量
(ModelScope / OpenAI / DashScope 等由环境变量选择)。
2026-04-10 16:52:07 +08:00
"""
import gc
2026-04-10 16:52:07 +08:00
import os
from dataclasses import dataclass
from typing import Any, List, Optional, Union
import numpy as np
import logging
2026-04-10 16:52:07 +08:00
logger = logging.getLogger(__name__)
@dataclass
class _RemoteEmbeddingEnv:
"""从环境变量解析出的远程 OpenAI-Compatible Embedding 配置。"""
api_key: str
base_url: str
model: str
max_batch: int
label: str
def _remote_embedding_from_env() -> _RemoteEmbeddingEnv:
"""
优先 ModelScope(MODELSCOPE_*);未配置时回退 DashScope(DASHSCOPE_*)。
"""
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
ms_base = os.getenv("MODELSCOPE_BASE_URL", "").strip()
ds_key = os.getenv("DASHSCOPE_API_KEY", "").strip()
ds_base = (
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
).strip()
# 仅当配置了 API Key 时走 ModelScope(避免仅有 BASE_URL 时误判、阻断 DashScope)
if ms_key:
base_url = ms_base or "https://api-inference.modelscope.cn/v1"
model = (
os.getenv("MODELSCOPE_EMBEDDING_MODEL")
or os.getenv("MODELSCOPE_MODEL", "Qwen/Qwen3-Embedding-8B")
).strip()
mb = os.getenv("MODELSCOPE_EMBEDDING_MAX_BATCH", "32").strip()
max_batch = max(1, int(mb)) if mb.isdigit() else 32
if not model:
raise ValueError("未配置 MODELSCOPE_EMBEDDING_MODEL(或 MODELSCOPE_MODEL)")
return _RemoteEmbeddingEnv(
api_key=ms_key,
base_url=base_url.rstrip("/"),
model=model,
max_batch=max_batch,
label="ModelScope",
)
# OpenAI 官方或兼容网关:OPENAI_API_KEY、OPENAI_EMBEDDING_MODEL、可选 OPENAI_BASE_URL
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
if oa_key:
if oa_key.startswith("http://") or oa_key.startswith("https://"):
raise ValueError(
"OPENAI_API_KEY 不能填写为 URL:请将网关地址写到 OPENAI_BASE_URL"
"(例如 http://host:9080/v1),密钥单独写在 OPENAI_API_KEY"
)
oa_base = (
os.getenv("OPENAI_BASE_URL", "").strip() or "https://api.openai.com/v1"
)
oa_model = os.getenv("OPENAI_EMBEDDING_MODEL", "").strip()
if not oa_model:
raise ValueError(
"使用 OpenAI 兼容 Embedding 时请设置 OPENAI_EMBEDDING_MODEL"
)
mb = os.getenv("OPENAI_EMBEDDING_MAX_BATCH", "100").strip()
max_batch = max(1, int(mb)) if mb.isdigit() else 100
return _RemoteEmbeddingEnv(
api_key=oa_key,
base_url=oa_base.rstrip("/"),
model=oa_model,
max_batch=max_batch,
label="OpenAI",
)
if not ds_key:
raise ValueError(
"远程 Embedding 未配置:请设置 MODELSCOPE_API_KEY(及可选 BASE_URL),"
"或 OPENAI_API_KEY / OPENAI_EMBEDDING_MODEL(及可选 OPENAI_BASE_URL),"
"或 DASHSCOPE_API_KEY / DASHSCOPE_BASE_URL / DASHSCOPE_MODEL"
)
if not ds_base:
raise ValueError("未配置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
model = os.getenv("DASHSCOPE_MODEL", "").strip()
if not model:
raise ValueError("未配置 DASHSCOPE_MODEL")
mb = os.getenv("DASHSCOPE_EMBEDDING_MAX_BATCH", "10").strip()
max_batch = max(1, int(mb)) if mb.isdigit() else 10
return _RemoteEmbeddingEnv(
api_key=ds_key,
base_url=ds_base.rstrip("/"),
model=model,
max_batch=max_batch,
label="DashScope",
)
class OpenAICompatibleRemoteEmbedding:
def __init__(
self,
api_key: Optional[str] = None,
base_url: Optional[str] = None,
model: Optional[str] = None,
max_batch: Optional[int] = None,
provider_label: Optional[str] = None,
):
try:
from openai import OpenAI
except ImportError as e:
raise ImportError(
"使用远程 Embedding 需要安装 openai:pip install openai"
) from e
if api_key is not None and base_url is not None and model is not None:
cfg = _RemoteEmbeddingEnv(
api_key=api_key.strip(),
base_url=base_url.strip().rstrip("/"),
model=model.strip(),
max_batch=max(1, int(max_batch)) if max_batch is not None else 32,
label=provider_label or "custom",
)
else:
cfg = _remote_embedding_from_env()
self.api_key = cfg.api_key
self.base_url = cfg.base_url
self.model = cfg.model
self._api_max_batch = cfg.max_batch
self._provider_label = cfg.label
self._client = OpenAI(api_key=self.api_key, base_url=self.base_url)
vd = os.getenv("VECTOR_DIM", "").strip()
self._embedding_dim: Optional[int] = int(vd) if vd.isdigit() else None
logger.info(
"使用 %s Embedding API:model=%s,base_url=%s,max_batch=%s",
self._provider_label,
self.model,
self.base_url,
self._api_max_batch,
)
@property
def embedding_dim(self) -> int:
if self._embedding_dim is None:
raise RuntimeError(
"尚未获知向量维度:请先执行一次 encode,或在 .env 中设置 VECTOR_DIM"
)
return self._embedding_dim
def _set_dim_from_vector(self, vec: List[float]) -> None:
if self._embedding_dim is None:
self._embedding_dim = len(vec)
logger.info("[OK] Embedding 向量维度:%s", self._embedding_dim)
def encode(
self,
texts: Union[str, List[str]],
batch_size: int = 10,
normalize: bool = True,
max_length: int = 8192,
show_progress: bool = False,
) -> np.ndarray:
del max_length # API 侧截断,此处仅保持签名与历史调用方一致
2026-04-10 16:52:07 +08:00
if isinstance(texts, str):
texts = [texts]
if not texts:
dim = self._embedding_dim
if dim is None:
vd = os.getenv("VECTOR_DIM", "").strip()
dim = int(vd) if vd.isdigit() else 1024
return np.empty((0, dim), dtype=np.float32)
# 无论调用方传多大,不能超过远端接口单次条数上限
step = max(1, min(int(batch_size), self._api_max_batch))
all_embeddings: List[np.ndarray] = []
iterator = range(0, len(texts), step)
if show_progress:
try:
from tqdm import tqdm
iterator = tqdm(iterator, desc="Embedding (API)")
except ImportError:
pass
for i in iterator:
batch = texts[i : i + step]
resp = self._client.embeddings.create(
model=self.model,
input=batch,
encoding_format="float",
)
rows = sorted(
[(d.index, d.embedding) for d in resp.data],
key=lambda x: x[0],
)
batch_embs = np.array([e for _, e in rows], dtype=np.float32)
if batch_embs.size > 0:
self._set_dim_from_vector(batch_embs[0].tolist())
if normalize:
norms = np.linalg.norm(batch_embs, axis=1, keepdims=True)
batch_embs = batch_embs / (norms + 1e-10)
all_embeddings.append(batch_embs)
return np.vstack(all_embeddings).astype(np.float32)
def similarity(self, emb1: np.ndarray, emb2: np.ndarray) -> np.ndarray:
return np.dot(emb1, emb2.T)
def encode_and_search(
self,
query: str,
documents: List[str],
top_k: int = 5,
) -> List[dict]:
query_emb = self.encode([query], normalize=True)
doc_embs = self.encode(documents, normalize=True)
scores = self.similarity(query_emb, doc_embs)[0]
top_indices = np.argsort(scores)[::-1][:top_k]
return [
{"score": float(scores[idx]), "document": documents[idx], "index": int(idx)}
for idx in top_indices
]
# 向后兼容旧名称
DashScopeOpenAIEmbedding = OpenAICompatibleRemoteEmbedding
_embedding_instance: Optional[Any] = None
def get_embedder(
_model_path: Optional[str] = None,
_device: Optional[str] = None,
2026-04-10 16:52:07 +08:00
force_reload: bool = False,
) -> OpenAICompatibleRemoteEmbedding:
2026-04-10 16:52:07 +08:00
"""
获取远程 Embedding 单例。``_model_path`` / ``_device`` 已废弃,仅为兼容旧调用保留。
2026-04-10 16:52:07 +08:00
"""
global _embedding_instance
if force_reload or _embedding_instance is None:
_embedding_instance = OpenAICompatibleRemoteEmbedding()
2026-04-10 16:52:07 +08:00
return _embedding_instance # type: ignore[return-value]
2026-04-10 16:52:07 +08:00
def clear_embedder() -> None:
"""清空单例(用于测试或切换远端配置)。"""
2026-04-10 16:52:07 +08:00
global _embedding_instance
_embedding_instance = None
gc.collect()