528 lines
18 KiB
Python
528 lines
18 KiB
Python
"""
|
||
Embedding 封装:兼容 OpenAI /v1/embeddings 的远程 API(默认),
|
||
或可选本地 HuggingFace 目录(USE_LOCAL_EMBEDDING=true + EMBEDDING_MODEL_PATH)。
|
||
"""
|
||
|
||
import os
|
||
from typing import Union, List, Optional, Any
|
||
import numpy as np
|
||
from pathlib import Path
|
||
|
||
try:
|
||
from transformers import AutoModel, AutoTokenizer
|
||
import torch
|
||
_TRANSFORMERS_AVAILABLE = True
|
||
except ImportError:
|
||
_TRANSFORMERS_AVAILABLE = False
|
||
|
||
import logging
|
||
from dataclasses import dataclass
|
||
|
||
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 Qwen3Embedding:
|
||
"""
|
||
本地 HuggingFace 格式 Embedding 模型(Mean Pooling + L2,用于向量检索)。
|
||
|
||
仅在 ``USE_LOCAL_EMBEDDING=true`` 时使用;路径由 ``model_path`` 或环境变量
|
||
``EMBEDDING_MODEL_PATH`` 指定,**不再内置默认目录**。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
model_path: Optional[str] = None,
|
||
device: Optional[str] = None,
|
||
use_fp16: bool = False
|
||
):
|
||
"""
|
||
初始化 embedding 模型
|
||
|
||
Args:
|
||
model_path: 本地模型目录;None 或空字符串时读 ``EMBEDDING_MODEL_PATH``
|
||
device: 推理设备('cpu', 'cuda', 'cuda:0'等),None则自动选择
|
||
use_fp16: 是否使用FP16混合精度(GPU可用时建议开启,速度更快)
|
||
"""
|
||
if not _TRANSFORMERS_AVAILABLE:
|
||
raise ImportError(
|
||
"transformers 和 torch 未安装。请运行:\n"
|
||
"pip install transformers torch sentencepiece accelerate"
|
||
)
|
||
|
||
resolved = (model_path or "").strip() or os.getenv("EMBEDDING_MODEL_PATH", "").strip()
|
||
if not resolved:
|
||
raise ValueError(
|
||
"已启用本地 Embedding(USE_LOCAL_EMBEDDING=true),但未设置有效模型路径。"
|
||
"请在 .env 中设置 EMBEDDING_MODEL_PATH 指向本地模型目录,"
|
||
"或设置 USE_LOCAL_EMBEDDING=false 使用 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程接口。"
|
||
)
|
||
|
||
model_path = Path(resolved)
|
||
if not model_path.exists():
|
||
raise FileNotFoundError(
|
||
f"本地 Embedding 模型目录不存在:{model_path}\n"
|
||
"请修正 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false。"
|
||
)
|
||
|
||
# 确定设备
|
||
if device is None:
|
||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||
|
||
self.device = device
|
||
logger.info("加载本地 Embedding 模型:%s,设备:%s", model_path, device)
|
||
|
||
# 加载 tokenizer:fast(Rust) 解析 tokenizer.json 需较新 tokenizers;
|
||
# 旧版本会报 ModelWrapper / untagged enum,回退到慢速 tokenizer 可恢复。
|
||
try:
|
||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||
str(model_path), trust_remote_code=True
|
||
)
|
||
except Exception as e:
|
||
err = str(e).lower()
|
||
if "modelwrapper" in err or "untagged enum" in err:
|
||
logger.warning(
|
||
"快速 tokenizer 解析 tokenizer.json 失败(多为 tokenizers 过旧),"
|
||
"改用 use_fast=False:%s",
|
||
e,
|
||
)
|
||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||
str(model_path), use_fast=False, trust_remote_code=True
|
||
)
|
||
else:
|
||
raise
|
||
|
||
try:
|
||
self.model = AutoModel.from_pretrained(
|
||
str(model_path), trust_remote_code=True
|
||
)
|
||
except ValueError as e:
|
||
msg = str(e)
|
||
if "qwen3" in msg.lower() or "does not recognize this architecture" in msg:
|
||
raise RuntimeError(
|
||
"当前 transformers 版本不支持 Qwen3(model_type=qwen3)。"
|
||
"请升级:pip install \"transformers>=4.51.0\" \"tokenizers>=0.21\""
|
||
) from e
|
||
raise
|
||
|
||
# 设置为评估模式并移动设备
|
||
self.model.eval()
|
||
self.model.to(device)
|
||
|
||
# 混合精度(仅GPU)
|
||
self.use_fp16 = use_fp16 and device != "cpu"
|
||
if self.use_fp16:
|
||
self.model.half()
|
||
|
||
# 嵌入维度
|
||
self.embedding_dim = self.model.config.hidden_size
|
||
logger.info(f"[OK] 模型加载完成,嵌入维度:{self.embedding_dim}")
|
||
|
||
def encode(
|
||
self,
|
||
texts: Union[str, List[str]],
|
||
batch_size: int = 32,
|
||
normalize: bool = True,
|
||
max_length: int = 8192,
|
||
show_progress: bool = False
|
||
) -> np.ndarray:
|
||
"""
|
||
编码文本为向量
|
||
|
||
Args:
|
||
texts: 单个文本或文本列表
|
||
batch_size: 批处理大小(根据显存调整)
|
||
normalize: 是否L2归一化(余弦相似度必需)
|
||
max_length: 最大序列长度(模型支持8192,建议512-1024平衡速度与精度)
|
||
show_progress: 是否显示进度条(需安装tqdm)
|
||
|
||
Returns:
|
||
numpy数组,shape=(len(texts), embedding_dim)
|
||
"""
|
||
if isinstance(texts, str):
|
||
texts = [texts]
|
||
|
||
if not texts:
|
||
return np.empty((0, self.embedding_dim), dtype=np.float32)
|
||
|
||
all_embeddings = []
|
||
|
||
# 可选进度条
|
||
iterator = range(0, len(texts), batch_size)
|
||
if show_progress:
|
||
try:
|
||
from tqdm import tqdm
|
||
iterator = tqdm(iterator, desc="Embedding")
|
||
except ImportError:
|
||
pass
|
||
|
||
for i in iterator:
|
||
batch = texts[i:i + batch_size]
|
||
|
||
# Tokenize
|
||
inputs = self.tokenizer(
|
||
batch,
|
||
padding=True,
|
||
truncation=True,
|
||
max_length=max_length,
|
||
return_tensors="pt"
|
||
).to(self.device)
|
||
|
||
# Inference
|
||
with torch.no_grad():
|
||
outputs = self.model(**inputs)
|
||
# Mean Pooling: 取序列维度的平均值
|
||
# outputs.last_hidden_state shape: (batch, seq_len, hidden_size)
|
||
embeddings = outputs.last_hidden_state.mean(dim=1)
|
||
|
||
# 转换为numpy(保持在CPU)
|
||
if self.device != "cpu":
|
||
embeddings = embeddings.cpu()
|
||
embeddings = embeddings.numpy()
|
||
|
||
if normalize:
|
||
# L2归一化(余弦相似度必需)
|
||
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
|
||
embeddings = embeddings / (norms + 1e-10)
|
||
|
||
all_embeddings.append(embeddings)
|
||
|
||
return np.vstack(all_embeddings).astype(np.float32)
|
||
|
||
def similarity(
|
||
self,
|
||
emb1: np.ndarray,
|
||
emb2: np.ndarray
|
||
) -> np.ndarray:
|
||
"""
|
||
计算两组embedding的余弦相似度
|
||
|
||
Args:
|
||
emb1: 第一组向量 (n, dim)
|
||
emb2: 第二组向量 (m, dim)
|
||
|
||
Returns:
|
||
相似度矩阵 (n, m),值域[-1, 1](若已归一化则为[0, 1])
|
||
"""
|
||
# 确保已归一化
|
||
return np.dot(emb1, emb2.T)
|
||
|
||
def encode_and_search(
|
||
self,
|
||
query: str,
|
||
documents: List[str],
|
||
top_k: int = 5
|
||
) -> List[dict]:
|
||
"""
|
||
便捷方法:编码查询并检索最相似的文档
|
||
|
||
Args:
|
||
query: 查询文本
|
||
documents: 候选文档列表
|
||
top_k: 返回前K个结果
|
||
|
||
Returns:
|
||
[{"score": float, "document": str, "index": int}, ...]
|
||
"""
|
||
query_emb = self.encode([query], normalize=True)
|
||
doc_embs = self.encode(documents, normalize=True)
|
||
|
||
scores = self.similarity(query_emb, doc_embs)[0]
|
||
|
||
# 获取top_k
|
||
top_indices = np.argsort(scores)[::-1][:top_k]
|
||
|
||
results = []
|
||
for idx in top_indices:
|
||
results.append({
|
||
"score": float(scores[idx]),
|
||
"document": documents[idx],
|
||
"index": int(idx)
|
||
})
|
||
|
||
return results
|
||
|
||
|
||
def _env_flag(name: str, default: str = "true") -> bool:
|
||
return os.getenv(name, default).strip().lower() in ("1", "true", "yes", "on")
|
||
|
||
|
||
class OpenAICompatibleRemoteEmbedding:
|
||
"""
|
||
通过 OpenAI 兼容接口获取文本向量(POST /v1/embeddings)。
|
||
|
||
环境变量(优先级:ModelScope → OpenAI 兼容 → DashScope):
|
||
- ModelScope:MODELSCOPE_API_KEY、可选 MODELSCOPE_BASE_URL(默认
|
||
https://api-inference.modelscope.cn/v1)、MODELSCOPE_EMBEDDING_MODEL
|
||
或 MODELSCOPE_MODEL、可选 MODELSCOPE_EMBEDDING_MAX_BATCH
|
||
- OpenAI 兼容:OPENAI_API_KEY、OPENAI_EMBEDDING_MODEL、可选 OPENAI_BASE_URL
|
||
(默认 https://api.openai.com/v1)、可选 OPENAI_EMBEDDING_MAX_BATCH
|
||
- DashScope:DASHSCOPE_API_KEY、DASHSCOPE_BASE_URL、DASHSCOPE_MODEL、
|
||
可选 DASHSCOPE_EMBEDDING_MAX_BATCH
|
||
|
||
可选 VECTOR_DIM:在首次请求前确定空列表返回的维度。
|
||
"""
|
||
|
||
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,
|
||
) -> Any:
|
||
"""
|
||
获取 Embedding 单例:默认 ``USE_LOCAL_EMBEDDING=false``,使用远程
|
||
OpenAI 兼容 / ModelScope / DashScope;为 true 时用本地目录(EMBEDDING_MODEL_PATH)。
|
||
"""
|
||
global _embedding_instance
|
||
|
||
if force_reload or _embedding_instance is None:
|
||
if _env_flag("USE_LOCAL_EMBEDDING", "false"):
|
||
_embedding_instance = Qwen3Embedding(
|
||
model_path=model_path,
|
||
device=device,
|
||
)
|
||
else:
|
||
_embedding_instance = OpenAICompatibleRemoteEmbedding()
|
||
|
||
return _embedding_instance
|
||
|
||
|
||
def clear_embedder():
|
||
"""清空单例(用于测试或切换模型)"""
|
||
global _embedding_instance
|
||
_embedding_instance = None
|
||
import gc
|
||
|
||
gc.collect()
|
||
if _TRANSFORMERS_AVAILABLE:
|
||
import torch
|
||
|
||
if torch.cuda.is_available():
|
||
torch.cuda.empty_cache()
|