267 lines
9.0 KiB
Python
267 lines
9.0 KiB
Python
"""
|
||
Embedding 封装:通过 OpenAI 兼容 ``/v1/embeddings`` 远程 API 获取向量
|
||
(ModelScope / OpenAI / DashScope 等由环境变量选择)。
|
||
"""
|
||
|
||
import gc
|
||
import os
|
||
from dataclasses import dataclass
|
||
from typing import Any, List, Optional, Union
|
||
|
||
import numpy as np
|
||
import logging
|
||
|
||
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 侧截断,此处仅保持签名与历史调用方一致
|
||
|
||
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,
|
||
force_reload: bool = False,
|
||
) -> OpenAICompatibleRemoteEmbedding:
|
||
"""
|
||
获取远程 Embedding 单例。``_model_path`` / ``_device`` 已废弃,仅为兼容旧调用保留。
|
||
"""
|
||
global _embedding_instance
|
||
|
||
if force_reload or _embedding_instance is None:
|
||
_embedding_instance = OpenAICompatibleRemoteEmbedding()
|
||
|
||
return _embedding_instance # type: ignore[return-value]
|
||
|
||
|
||
def clear_embedder() -> None:
|
||
"""清空单例(用于测试或切换远端配置)。"""
|
||
global _embedding_instance
|
||
_embedding_instance = None
|
||
gc.collect()
|