Files
ai-g3sb-backman2.0/backend/utils/embedding.py
T

267 lines
9.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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()