Refactor embedding configuration to remove local model support, transitioning to a unified remote API approach. Update environment variables and documentation accordingly. Enhance error handling in the orchestrator and related modules to reflect these changes. This update simplifies the embedding process and improves overall system reliability.
This commit is contained in:
+16
-277
@@ -1,22 +1,15 @@
|
||||
"""
|
||||
Embedding 封装:兼容 OpenAI /v1/embeddings 的远程 API(默认),
|
||||
或可选本地 HuggingFace 目录(USE_LOCAL_EMBEDDING=true + EMBEDDING_MODEL_PATH)。
|
||||
Embedding 封装:通过 OpenAI 兼容 ``/v1/embeddings`` 远程 API 获取向量
|
||||
(ModelScope / OpenAI / DashScope 等由环境变量选择)。
|
||||
"""
|
||||
|
||||
import gc
|
||||
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
|
||||
from typing import Any, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -110,246 +103,7 @@ def _remote_embedding_from_env() -> _RemoteEmbeddingEnv:
|
||||
)
|
||||
|
||||
|
||||
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,
|
||||
@@ -416,7 +170,7 @@ class OpenAICompatibleRemoteEmbedding:
|
||||
max_length: int = 8192,
|
||||
show_progress: bool = False,
|
||||
) -> np.ndarray:
|
||||
del max_length # API 侧截断,此处仅保持签名与本地实现一致
|
||||
del max_length # API 侧截断,此处仅保持签名与历史调用方一致
|
||||
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
@@ -486,42 +240,27 @@ class OpenAICompatibleRemoteEmbedding:
|
||||
# 向后兼容旧名称
|
||||
DashScopeOpenAIEmbedding = OpenAICompatibleRemoteEmbedding
|
||||
|
||||
# 全局单例(避免重复加载模型,节省显存/内存)
|
||||
_embedding_instance: Optional[Any] = None
|
||||
|
||||
|
||||
def get_embedder(
|
||||
model_path: Optional[str] = None,
|
||||
device: Optional[str] = None,
|
||||
_model_path: Optional[str] = None,
|
||||
_device: Optional[str] = None,
|
||||
force_reload: bool = False,
|
||||
) -> Any:
|
||||
) -> OpenAICompatibleRemoteEmbedding:
|
||||
"""
|
||||
获取 Embedding 单例:默认 ``USE_LOCAL_EMBEDDING=false``,使用远程
|
||||
OpenAI 兼容 / ModelScope / DashScope;为 true 时用本地目录(EMBEDDING_MODEL_PATH)。
|
||||
获取远程 Embedding 单例。``_model_path`` / ``_device`` 已废弃,仅为兼容旧调用保留。
|
||||
"""
|
||||
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()
|
||||
_embedding_instance = OpenAICompatibleRemoteEmbedding()
|
||||
|
||||
return _embedding_instance
|
||||
return _embedding_instance # type: ignore[return-value]
|
||||
|
||||
|
||||
def clear_embedder():
|
||||
"""清空单例(用于测试或切换模型)"""
|
||||
def clear_embedder() -> None:
|
||||
"""清空单例(用于测试或切换远端配置)。"""
|
||||
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