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:
陈辅元
2026-04-15 09:49:18 +08:00
parent 4ea3056e95
commit 2d0b1ab36f
21 changed files with 103 additions and 405 deletions
Binary file not shown.
+2 -8
View File
@@ -46,7 +46,6 @@ class Text2SQLOrchestrator:
schema_manager: SchemaManager,
deepseek_api_key: Optional[str] = None,
deepseek_config: Optional[DeepSeekConfig] = None,
embedding_model_path: Optional[str] = None,
vector_db_path: str = "./data/embeddings/chroma",
max_retry: int = 2,
use_vector_search: bool = True,
@@ -64,7 +63,6 @@ class Text2SQLOrchestrator:
schema_manager: Schema管理器实例
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
deepseek_config: DeepSeek配置对象(优先于api_key)
embedding_model_path: 本地 Embedding 模型目录(仅 USE_LOCAL_EMBEDDING=true)
vector_db_path: 向量数据库路径
max_retry: 最大重试次数(包含首次生成)
use_vector_search: 是否使用向量检索粗筛
@@ -86,7 +84,6 @@ class Text2SQLOrchestrator:
# 初始化向量索引(延迟加载)
self._vector_index: Optional[SchemaIndexer] = None
self._vector_db_path = vector_db_path
self._embedding_model_path = embedding_model_path
# Few-shot 初始化
self.fewshot_enabled = fewshot_enabled
@@ -110,10 +107,7 @@ class Text2SQLOrchestrator:
else:
# Chroma 优先时默认不再依赖 JSONL;否则保留原默认路径
path = "" if use_chroma else "./data/experiences/all_samples.jsonl"
self.fewshot_selector = FewShotSelector(
path or None,
embedding_model_path=self._embedding_model_path,
)
self.fewshot_selector = FewShotSelector(path or None)
logger.info(
f"Few-shot已启用: top_k={fewshot_top_k}, "
f"min_rating={fewshot_min_rating}"
@@ -151,7 +145,7 @@ class Text2SQLOrchestrator:
if self._vector_index is None:
from utils.embedding import get_embedder
embedder = get_embedder(self._embedding_model_path)
embedder = get_embedder()
self._vector_index = SchemaIndexer(
embedder=embedder,
persist_dir=self._vector_db_path
Binary file not shown.
+1 -3
View File
@@ -16,10 +16,8 @@ class Settings(BaseSettings):
temperature: float = 0.3
max_tokens: int = 4096
# Embedding配置
embedding_model_path: str = "" # 仅本地 Embedding 时使用;默认走远程 API
# Embedding / 向量检索(编码走远程 API,见 OPENAI_* / MODELSCOPE_* 等)
vector_dim: int = 2048
use_local_embedding: bool = True
# Schema配置
schema_dir: str = "./data/schemas"
+24 -54
View File
@@ -37,11 +37,6 @@ def _load_project_env():
load_dotenv(_repo_root() / ".env")
def _embedding_model_path() -> str:
"""仅 ``USE_LOCAL_EMBEDDING=true`` 时需要;远程 Embedding 可为空。"""
return os.getenv("EMBEDDING_MODEL_PATH", "").strip()
def resolve_sql_dialect(name: str) -> str:
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
n = (name or "sqlserver").lower().strip()
@@ -53,52 +48,33 @@ def resolve_sql_dialect(name: str) -> str:
def setup_environment():
"""环境检查"""
_load_project_env()
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in (
"1", "true", "yes", "on",
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
oa_key_ok = oa_key and not (
oa_key.startswith("http://") or oa_key.startswith("https://")
)
if use_local_emb:
mp_str = _embedding_model_path()
if not mp_str:
logger.warning(
"USE_LOCAL_EMBEDDING=true 但未设置 EMBEDDING_MODEL_PATH(本地模型目录)"
)
return False
model_path = Path(mp_str)
if not model_path.exists():
logger.warning("本地 Embedding 目录不存在: %s", model_path)
logger.info(
"请设置正确的 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false,"
"并配置 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding"
)
if ms_key:
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
elif oa_key_ok:
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
logger.warning("未设置 OPENAI_EMBEDDING_MODEL(远程 Embedding 模型名)")
return False
else:
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
oa_key_ok = oa_key and not (
oa_key.startswith("http://") or oa_key.startswith("https://")
)
if ms_key:
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
elif oa_key_ok:
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
logger.warning("USE_LOCAL_EMBEDDING=false 但未设置 OPENAI_EMBEDDING_MODEL")
return False
else:
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
logger.warning(
"USE_LOCAL_EMBEDDING=false 但未设置 MODELSCOPE_API_KEY、"
"OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
)
return False
base = (
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
).strip()
if not base:
logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
return False
if not os.getenv("DASHSCOPE_MODEL", "").strip():
logger.warning("未设置 DASHSCOPE_MODEL")
return False
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
logger.warning(
"未设置 MODELSCOPE_API_KEY、"
"OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
)
return False
base = (
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
).strip()
if not base:
logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
return False
if not os.getenv("DASHSCOPE_MODEL", "").strip():
logger.warning("未设置 DASHSCOPE_MODEL")
return False
# 检查Schema文件(支持相对路径和绝对路径)
schema_path_str = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
@@ -214,7 +190,6 @@ def create_orchestrator(schema_mgr, args):
orchestrator = Text2SQLOrchestrator(
schema_manager=schema_mgr,
deepseek_config=config,
embedding_model_path=args.embedding_model,
vector_db_path=args.vector_db,
max_retry=args.max_retry,
use_vector_search=not args.no_vector_search,
@@ -366,11 +341,6 @@ def main():
help="最大重试次数(默认: 2)"
)
_load_project_env()
parser.add_argument(
"--embedding-model",
default=_embedding_model_path(),
help="Embedding模型路径(默认来自环境变量 EMBEDDING_MODEL_PATH)"
)
parser.add_argument(
"--vector-db",
default=os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma").strip(),
Binary file not shown.
+16 -277
View File
@@ -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()
+4 -16
View File
@@ -1,8 +1,8 @@
"""
Few-shot示例选择器 - 基于经验数据集动态选择相关示例
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(本地 Qwen3 或远程 OpenAI 兼容 API,
由 USE_LOCAL_EMBEDDING 及 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(远程 OpenAI 兼容 API;
由 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
用法:
from utils.fewshot_selector import FewShotSelector
@@ -26,10 +26,6 @@ import logging
logger = logging.getLogger(__name__)
# 仅 USE_LOCAL_EMBEDDING=true 时通过 EMBEDDING_MODEL_PATH 使用;远程模式留空即可
_DEFAULT_LOCAL_EMBED_PATH = ""
@dataclass
class ExperienceSample:
"""经验数据样本"""
@@ -84,7 +80,6 @@ class FewShotSelector:
def __init__(
self,
samples_path: Optional[str] = None,
embedding_model_path: Optional[str] = None,
use_cache: bool = True,
use_chroma: Optional[bool] = None,
chroma_persist_dir: Optional[str] = None,
@@ -93,8 +88,6 @@ class FewShotSelector:
Args:
samples_path: 样本 JSONL 路径;**空字符串**表示不读文件(仅当 ``FEWSHOT_USE_CHROMA=true``
且 Chroma 中已有数据时可用)。构建索引仍请用 ``scripts/build_fewshot_chroma_index.py --samples``。
embedding_model_path: 本地 Embedding 模型目录;None 时用环境变量
EMBEDDING_MODEL_PATH(仅 USE_LOCAL_EMBEDDING=true 时有效)
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
use_chroma: 是否使用 Chroma 持久化向量库;None 时读环境变量 FEWSHOT_USE_CHROMA
chroma_persist_dir: Chroma 目录;None 时用 FEWSHOT_CHROMA_PATH 或默认 chroma_fewshot
@@ -113,11 +106,6 @@ class FewShotSelector:
else os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in ("1", "true", "yes")
)
self._chroma_persist_dir = chroma_persist_dir
self._embedding_model_path = (
embedding_model_path
if embedding_model_path is not None
else os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_LOCAL_EMBED_PATH).strip()
)
if self.use_chroma:
self._init_chroma_mode()
@@ -130,7 +118,7 @@ class FewShotSelector:
from utils.embedding import get_embedder
from utils.fewshot_chroma_store import FewShotChromaStore
self._embedder = get_embedder(self._embedding_model_path)
self._embedder = get_embedder()
self._chroma_store = FewShotChromaStore(
self._embedder,
persist_dir=self._chroma_persist_dir,
@@ -199,7 +187,7 @@ class FewShotSelector:
"""非 Chroma:初始化 Embedder 与内存 numpy 索引。"""
from utils.embedding import get_embedder
self._embedder = get_embedder(self._embedding_model_path)
self._embedder = get_embedder()
self._build_numpy_index()
def _build_numpy_index(self) -> None: