Files
ai-g3sb-backman2.0/utils/embedding.py
T
2026-04-10 16:52:07 +08:00

531 lines
18 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 封装:本地 Qwen3-Embedding,或兼容 OpenAI /v1/embeddings 的远程 API
(ModelScope 推理、阿里云 DashScope 等)。
"""
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:
"""
Qwen3-Embedding-0.6B 向量化封装
使用 Mean Pooling 将token embeddings聚合为句子向量,
并进行L2归一化以支持余弦相似度计算。
"""
def __init__(
self,
model_path: Optional[str] = None,
device: Optional[str] = None,
use_fp16: bool = False
):
"""
初始化 embedding 模型
Args:
model_path: 本地模型路径,若为None则从环境变量或默认路径加载
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"
)
# 确定模型路径
if model_path is None:
model_path = os.getenv(
"EMBEDDING_MODEL_PATH",
"./data/models/Qwen3-Embedding-0.6B"
)
model_path = Path(model_path)
if not model_path.exists():
raise FileNotFoundError(
f"模型目录不存在:{model_path}\n"
"请先下载模型:\n"
" modelscope download --model 'Qwen/Qwen3-Embedding-0.6B' "
f"--local_dir '{model_path}'\n"
"或从Hugging Face下载:git lfs install && git clone "
f"https://huggingface.co/Qwen/Qwen3-Embedding-0.6B {model_path}"
)
# 确定设备
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
self.device = device
logger.info(f"加载Qwen3-Embedding模型:{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=true 时用本地 Qwen3,否则用远程
OpenAI 兼容 API(优先级见 OpenAICompatibleRemoteEmbedding)。
"""
global _embedding_instance
if force_reload or _embedding_instance is None:
if _env_flag("USE_LOCAL_EMBEDDING", "true"):
_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()