0.1.1 暂存
This commit is contained in:
@@ -0,0 +1,530 @@
|
||||
"""
|
||||
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()
|
||||
Reference in New Issue
Block a user