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
+1 -4
View File
@@ -13,10 +13,7 @@ MAX_TOKENS=4096
# 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等 # 生成与校验默认 SQL Server(T-SQL);可改为 mysql / postgresql 等
TEXT2SQL_DIALECT=sqlserver TEXT2SQL_DIALECT=sqlserver
# ========== Embedding 配置 ========== # ========== Embedding 配置(远程 /v1/embeddings,须配置下方 OPENAI_* 或 MODELSCOPE_* 等)==========
# 默认 false:使用远程 /v1/embeddings(须配置下方 OPENAI_* 或 MODELSCOPE_* 等)
# 仅当 true 时需设置 EMBEDDING_MODEL_PATH 指向本地 HuggingFace 模型目录
USE_LOCAL_EMBEDDING=false
# ModelScope 推理 API:https://www.modelscope.cn/docs/inference-api # ModelScope 推理 API:https://www.modelscope.cn/docs/inference-api
# MODELSCOPE_API_KEY=ms-b95f3233-4797-4efd-90e6-b3728743571d # MODELSCOPE_API_KEY=ms-b95f3233-4797-4efd-90e6-b3728743571d
# MODELSCOPE_BASE_URL=https://api-inference.modelscope.cn/v1 # MODELSCOPE_BASE_URL=https://api-inference.modelscope.cn/v1
+44
View File
@@ -162,3 +162,47 @@
| 配置项 | 含义 | 默认 | | 配置项 | 含义 | 默认 |
|--------|------|------| |--------|------|------|
| `DIALOG_INTENT_CLASSIFIER` | `hybrid`:规则快速路径 + LLM;`rules`:仅规则 | `hybrid`(未设置时按 hybrid 处理) | | `DIALOG_INTENT_CLASSIFIER` | `hybrid`:规则快速路径 + LLM;`rules`:仅规则 | `hybrid`(未设置时按 hybrid 处理) |
---
# Impact Analysis — 移除本地 Embedding(仅远程 API)
## 1. 改动概览
- **背景与目标**:不再维护本地 HuggingFace / PyTorch 推理路径;Embedding 统一为 OpenAI 兼容远程 `/v1/embeddings`。
- **涉及模块**:`backend/utils/embedding.py`、`backend/main.py`、`backend/agents/orchestrator.py`、`backend/utils/fewshot_selector.py`、`backend/config/settings.py`、`api_server.py`、`scripts/build_fewshot_chroma_index.py`、`pyproject.toml`、`requirements.txt`、`main.spec`、`README.md`、`.env` 注释。
- **改动类型**:功能删减 / 依赖精简。
## 2. 方法级改动
| 位置 | 变更 |
|------|------|
| `Qwen3Embedding`(原) | **删除**;本地 `transformers`+`torch` 编码路径移除。 |
| `get_embedder` | 仅构造 `OpenAICompatibleRemoteEmbedding`;前两个位置参数废弃保留以兼容旧调用。 |
| `clear_embedder` | 仅 `gc.collect()`,不再触碰 CUDA。 |
| `Text2SQLOrchestrator.__init__` | 移除 `embedding_model_path`;`FewShotSelector` 不再传模型路径。 |
| `FewShotSelector.__init__` | 移除 `embedding_model_path`。 |
| `setup_environment` | 删除 `USE_LOCAL_EMBEDDING` / `EMBEDDING_MODEL_PATH` 分支。 |
## 3. 调用方与影响范围
- **调用方**:所有原 `get_embedder(path)` 仍可运行(路径被忽略);`Text2SQLOrchestrator(..., embedding_model_path=...)` 需改为不传该参数(已改仓库内引用)。
- **破坏性变更**:**是**——不再支持 `USE_LOCAL_EMBEDDING=true` 与本地模型目录;`pyproject.toml` / `requirements.txt` 不再声明 `torch`/`transformers`/`sentencepiece`/`accelerate`/`safetensors`(若 `camel-ai[all]` 等仍带入部分传递依赖,以实际 lock 为准)。
## 4. 风险与回滚
- **风险级别**:中(仅使用本地 Embedding 的部署将失效;需改为远程 Key 与模型名)。
- **回滚**:回退提交并恢复依赖与 `embedding.py` 历史版本。
**回滚方式是否简单**:是(单提交回退)。
## 5. 验证与测试
- 建议:`python -m py_compile` 对改动 `.py`;在已配置 `OPENAI_*` 或 `MODELSCOPE_*` 的环境下跑一次 Schema 向量构建或 Few-shot 检索。
## 6. 配置变更
| 移除/失效项 | 说明 |
|-------------|------|
| `USE_LOCAL_EMBEDDING`、`EMBEDDING_MODEL_PATH` | 代码不再读取;`.env` 中可删去以免误解。 |
| CLI `--embedding-model` | 已删除。 |
+7 -21
View File
@@ -1,4 +1,3 @@
<<<<<<< HEAD
# Text2SQL 多智能体系统 # Text2SQL 多智能体系统
基于 **CAMEL AI 框架**、**DeepSeek 大模型** 和 **经验数据集 Few-Shot 增强** 的专业 Text-to-SQL 生成系统,专为证券经纪业务领域优化。 基于 **CAMEL AI 框架**、**DeepSeek 大模型** 和 **经验数据集 Few-Shot 增强** 的专业 Text-to-SQL 生成系统,专为证券经纪业务领域优化。
@@ -6,7 +5,7 @@
## 🚀 核心特性 ## 🚀 核心特性
- **🤖 多Agent协作流水线**:Schema Linker → SQL Generator → Validator 三阶段协作 - **🤖 多Agent协作流水线**:Schema Linker → SQL Generator → Validator 三阶段协作
- **🎯 高精度表检索**:向量检索(OpenAI 兼容 / ModelScope 等远程 Embedding,或可选本地模型)+ LLM精筛,从 250+ 张表中精准定位相关表 - **🎯 高精度表检索**:向量检索(OpenAI 兼容 / ModelScope 等远程 Embedding)+ LLM精筛,从 250+ 张表中精准定位相关表
- **💡 经验数据集 Few-Shot**:基于 50 条高质量样例的语义检索,动态注入相似示例提升准确率 - **💡 经验数据集 Few-Shot**:基于 50 条高质量样例的语义检索,动态注入相似示例提升准确率
- **🔒 双重验证机制**:程序语法验证 + LLM语义验证,确保 SQL 正确性 - **🔒 双重验证机制**:程序语法验证 + LLM语义验证,确保 SQL 正确性
- **🏦 金融领域深度优化**:针对证券经纪业务(账户、持仓、现金、结算)定制 Prompt 和推断规则 - **🏦 金融领域深度优化**:针对证券经纪业务(账户、持仓、现金、结算)定制 Prompt 和推断规则
@@ -85,7 +84,7 @@
|------|---------|-----------| |------|---------|-----------|
| Agent 框架 | **CAMEL AI** | >= 0.2.0,角色扮演、消息通信 | | Agent 框架 | **CAMEL AI** | >= 0.2.0,角色扮演、消息通信 |
| LLM 模型 | **DeepSeek-chat** | 国产大模型,代码生成能力强 | | LLM 模型 | **DeepSeek-chat** | 国产大模型,代码生成能力强 |
| Embedding | **远程 API(默认)** | OpenAI 兼容 / ModelScope 等;可选本地 HF 目录 | | Embedding | **远程 API** | OpenAI 兼容 / ModelScope / DashScope 等 |
| 向量数据库 | **ChromaDB** | 本地轻量,持久化 Schema 索引 | | 向量数据库 | **ChromaDB** | 本地轻量,持久化 Schema 索引 |
| SQL 解析 | **sqlglot** | >= 20.0.0,多方言 AST 转换与验证 | | SQL 解析 | **sqlglot** | >= 20.0.0,多方言 AST 转换与验证 |
| 配置管理 | **Pydantic Settings** | 类型安全的环境变量管理 | | 配置管理 | **Pydantic Settings** | 类型安全的环境变量管理 |
@@ -115,7 +114,7 @@ text2sql_agent_camel/
│ ├── llm/ │ ├── llm/
│ │ └── deepseek_client.py # DeepSeek API 客户端(chat / validate_sql / select_tables) │ │ └── deepseek_client.py # DeepSeek API 客户端(chat / validate_sql / select_tables)
│ └── utils/ │ └── utils/
│ ├── embedding.py # Embedding(远程 OpenAI 兼容 / 可选本地 HF) │ ├── embedding.py # Embedding(远程 OpenAI 兼容 /v1/embeddings)
│ ├── sql_parser.py # sqlglot 工具(语法验证、方言转换、规范化) │ ├── sql_parser.py # sqlglot 工具(语法验证、方言转换、规范化)
│ ├── validators.py # 验证逻辑(危险操作检测、完整验证流水线) │ ├── validators.py # 验证逻辑(危险操作检测、完整验证流水线)
│ └── fewshot_selector.py # Few-Shot 示例选择器(语义/关键词检索) │ └── fewshot_selector.py # Few-Shot 示例选择器(语义/关键词检索)
@@ -130,7 +129,7 @@ text2sql_agent_camel/
│ │ ├── fewshot_examples.md # Markdown 格式示例 │ │ ├── fewshot_examples.md # Markdown 格式示例
│ │ └── by_tag/ # 按标签分类(20 类) │ │ └── by_tag/ # 按标签分类(20 类)
│ ├── embeddings/ # ChromaDB 向量索引持久化目录 │ ├── embeddings/ # ChromaDB 向量索引持久化目录
│ └── models/ # 本地 Embedding 模型权重 │ └── models/ # (可选)其他本地模型资源
├── scripts/ ├── scripts/
│ ├── parse_examples.py # 解析 Example_text2sql.md → 结构化数据集 │ ├── parse_examples.py # 解析 Example_text2sql.md → 结构化数据集
│ ├── integrate_fewshot.py # Few-Shot 集成指南与测试 │ ├── integrate_fewshot.py # Few-Shot 集成指南与测试
@@ -149,7 +148,7 @@ text2sql_agent_camel/
- Python **3.10+** - Python **3.10+**
- DeepSeek API Key(申请地址:https://platform.deepseek.com/api_keys) - DeepSeek API Key(申请地址:https://platform.deepseek.com/api_keys)
- (可选)本地 Embedding 模型权重(或配置远程 Embedding API) - 配置远程 Embedding API(OPENAI_* / MODELSCOPE_* / DASHSCOPE_*)
### 1. 安装依赖 ### 1. 安装依赖
@@ -186,8 +185,7 @@ cp .env .env.local # 或直接编辑 .env
DEEPSEEK_API_KEY=sk-xxxxxxxxxxxxxxxxxxxxxxxx DEEPSEEK_API_KEY=sk-xxxxxxxxxxxxxxxxxxxxxxxx
DEEPSEEK_BASE_URL=https://api.deepseek.com DEEPSEEK_BASE_URL=https://api.deepseek.com
# 推荐:远程 Embedding(OpenAI 兼容网关或 ModelScope 等,无需本地模型目录) # 远程 Embedding(OpenAI 兼容网关或 ModelScope 等)
USE_LOCAL_EMBEDDING=false
OPENAI_BASE_URL=https://api.openai.com/v1 OPENAI_BASE_URL=https://api.openai.com/v1
OPENAI_API_KEY=sk-xxxxxxxx OPENAI_API_KEY=sk-xxxxxxxx
OPENAI_EMBEDDING_MODEL=text-embedding-3-small OPENAI_EMBEDDING_MODEL=text-embedding-3-small
@@ -195,10 +193,6 @@ OPENAI_EMBEDDING_MODEL=text-embedding-3-small
# 或使用 ModelScope # 或使用 ModelScope
# MODELSCOPE_API_KEY=ms-xxxxxxxx # MODELSCOPE_API_KEY=ms-xxxxxxxx
# MODELSCOPE_EMBEDDING_MODEL=Qwen/Qwen3-Embedding-8B # MODELSCOPE_EMBEDDING_MODEL=Qwen/Qwen3-Embedding-8B
# 可选:本地 HuggingFace 模型目录(需 USE_LOCAL_EMBEDDING=true + EMBEDDING_MODEL_PATH)
# USE_LOCAL_EMBEDDING=true
# EMBEDDING_MODEL_PATH=/path/to/local/embedding-model
``` ```
### 3. 准备 Schema 文件 ### 3. 准备 Schema 文件
@@ -359,13 +353,11 @@ FEWSHOT_USE_CHROMA=true
FEWSHOT_CHROMA_PATH=./data/embeddings/chroma_fewshot FEWSHOT_CHROMA_PATH=./data/embeddings/chroma_fewshot
# FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl # 仅构建索引时需要 # FEWSHOT_DATA_PATH=./data/experiences/all_samples.jsonl # 仅构建索引时需要
# ========== Embedding ========== # ========== Embedding(远程 /v1/embeddings)==========
USE_LOCAL_EMBEDDING=false # false=远程(默认);true 时需配置 EMBEDDING_MODEL_PATH
# OPENAI_* 或 MODELSCOPE_* / DASHSCOPE_*(与代码内优先级一致) # OPENAI_* 或 MODELSCOPE_* / DASHSCOPE_*(与代码内优先级一致)
OPENAI_BASE_URL=https://api.openai.com/v1 OPENAI_BASE_URL=https://api.openai.com/v1
OPENAI_API_KEY=sk-xxxx OPENAI_API_KEY=sk-xxxx
OPENAI_EMBEDDING_MODEL=text-embedding-3-small OPENAI_EMBEDDING_MODEL=text-embedding-3-small
# EMBEDDING_MODEL_PATH= # 仅 USE_LOCAL_EMBEDDING=true 时填写本地目录
# ========== Schema ========== # ========== Schema ==========
SCHEMA_DIR=./data/schemas SCHEMA_DIR=./data/schemas
@@ -397,7 +389,6 @@ python backend/main.py [选项]
# 向量检索 # 向量检索
--no-vector-search 禁用向量检索(使用全部表) --no-vector-search 禁用向量检索(使用全部表)
--embedding-model PATH Embedding 模型路径
# 其他 # 其他
--verbose, -v 详细日志 --verbose, -v 详细日志
@@ -741,8 +732,3 @@ MIT License — 详见 [LICENSE](LICENSE) 文件。
**维护者**:Text2SQL Team **维护者**:Text2SQL Team
**更新日期**:2026-04-10 **更新日期**:2026-04-10
**文档版本**:v0.2.0 **文档版本**:v0.2.0
=======
# ai-g3sb-backman2.0
ai-g3sb-backman2.0
>>>>>>> 267151bdd466eadcfdcb4b55256bf7a15466e466
Binary file not shown.
+1 -3
View File
@@ -98,8 +98,7 @@ def get_orchestrator():
if not setup_environment(): if not setup_environment():
raise RuntimeError( raise RuntimeError(
"环境检查失败(请查看上方 WARNING/INFO,常见原因:" "环境检查失败(请查看上方 WARNING/INFO,常见原因:"
"USE_LOCAL_EMBEDDING=true 但未配置或找不到 EMBEDDING_MODEL_PATH;" "未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding;"
"或 USE_LOCAL_EMBEDDING=false 但未配好 OPENAI_* / MODELSCOPE_* / DASHSCOPE_*;"
"或 Schema 文件路径不对、DEEPSEEK_API_KEY 未设置)" "或 Schema 文件路径不对、DEEPSEEK_API_KEY 未设置)"
) )
@@ -114,7 +113,6 @@ def get_orchestrator():
temperature = float(os.getenv("TEMPERATURE", "0.3")) temperature = float(os.getenv("TEMPERATURE", "0.3"))
max_tokens = int(os.getenv("MAX_TOKENS", "4096")) max_tokens = int(os.getenv("MAX_TOKENS", "4096"))
max_retry = int(os.getenv("MAX_RETRY", "2")) max_retry = int(os.getenv("MAX_RETRY", "2"))
embedding_model = None
vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma") vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma")
no_vector_search = False # 启用向量搜索(ChromaDB 已修复) no_vector_search = False # 启用向量搜索(ChromaDB 已修复)
no_fewshot = not os.getenv("FEWSHOT_ENABLED", "true").lower() in ("true", "1", "yes") no_fewshot = not os.getenv("FEWSHOT_ENABLED", "true").lower() in ("true", "1", "yes")
Binary file not shown.
+2 -8
View File
@@ -46,7 +46,6 @@ class Text2SQLOrchestrator:
schema_manager: SchemaManager, schema_manager: SchemaManager,
deepseek_api_key: Optional[str] = None, deepseek_api_key: Optional[str] = None,
deepseek_config: Optional[DeepSeekConfig] = None, deepseek_config: Optional[DeepSeekConfig] = None,
embedding_model_path: Optional[str] = None,
vector_db_path: str = "./data/embeddings/chroma", vector_db_path: str = "./data/embeddings/chroma",
max_retry: int = 2, max_retry: int = 2,
use_vector_search: bool = True, use_vector_search: bool = True,
@@ -64,7 +63,6 @@ class Text2SQLOrchestrator:
schema_manager: Schema管理器实例 schema_manager: Schema管理器实例
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY) deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
deepseek_config: DeepSeek配置对象(优先于api_key) deepseek_config: DeepSeek配置对象(优先于api_key)
embedding_model_path: 本地 Embedding 模型目录(仅 USE_LOCAL_EMBEDDING=true)
vector_db_path: 向量数据库路径 vector_db_path: 向量数据库路径
max_retry: 最大重试次数(包含首次生成) max_retry: 最大重试次数(包含首次生成)
use_vector_search: 是否使用向量检索粗筛 use_vector_search: 是否使用向量检索粗筛
@@ -86,7 +84,6 @@ class Text2SQLOrchestrator:
# 初始化向量索引(延迟加载) # 初始化向量索引(延迟加载)
self._vector_index: Optional[SchemaIndexer] = None self._vector_index: Optional[SchemaIndexer] = None
self._vector_db_path = vector_db_path self._vector_db_path = vector_db_path
self._embedding_model_path = embedding_model_path
# Few-shot 初始化 # Few-shot 初始化
self.fewshot_enabled = fewshot_enabled self.fewshot_enabled = fewshot_enabled
@@ -110,10 +107,7 @@ class Text2SQLOrchestrator:
else: else:
# Chroma 优先时默认不再依赖 JSONL;否则保留原默认路径 # Chroma 优先时默认不再依赖 JSONL;否则保留原默认路径
path = "" if use_chroma else "./data/experiences/all_samples.jsonl" path = "" if use_chroma else "./data/experiences/all_samples.jsonl"
self.fewshot_selector = FewShotSelector( self.fewshot_selector = FewShotSelector(path or None)
path or None,
embedding_model_path=self._embedding_model_path,
)
logger.info( logger.info(
f"Few-shot已启用: top_k={fewshot_top_k}, " f"Few-shot已启用: top_k={fewshot_top_k}, "
f"min_rating={fewshot_min_rating}" f"min_rating={fewshot_min_rating}"
@@ -151,7 +145,7 @@ class Text2SQLOrchestrator:
if self._vector_index is None: if self._vector_index is None:
from utils.embedding import get_embedder from utils.embedding import get_embedder
embedder = get_embedder(self._embedding_model_path) embedder = get_embedder()
self._vector_index = SchemaIndexer( self._vector_index = SchemaIndexer(
embedder=embedder, embedder=embedder,
persist_dir=self._vector_db_path 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 temperature: float = 0.3
max_tokens: int = 4096 max_tokens: int = 4096
# Embedding配置 # Embedding / 向量检索(编码走远程 API,见 OPENAI_* / MODELSCOPE_* 等)
embedding_model_path: str = "" # 仅本地 Embedding 时使用;默认走远程 API
vector_dim: int = 2048 vector_dim: int = 2048
use_local_embedding: bool = True
# Schema配置 # Schema配置
schema_dir: str = "./data/schemas" schema_dir: str = "./data/schemas"
+24 -54
View File
@@ -37,11 +37,6 @@ def _load_project_env():
load_dotenv(_repo_root() / ".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: def resolve_sql_dialect(name: str) -> str:
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。""" """CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
n = (name or "sqlserver").lower().strip() n = (name or "sqlserver").lower().strip()
@@ -53,52 +48,33 @@ def resolve_sql_dialect(name: str) -> str:
def setup_environment(): def setup_environment():
"""环境检查""" """环境检查"""
_load_project_env() _load_project_env()
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in ( ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
"1", "true", "yes", "on", 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: if ms_key:
mp_str = _embedding_model_path() pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
if not mp_str: elif oa_key_ok:
logger.warning( if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
"USE_LOCAL_EMBEDDING=true 但未设置 EMBEDDING_MODEL_PATH(本地模型目录)" logger.warning("未设置 OPENAI_EMBEDDING_MODEL(远程 Embedding 模型名)")
)
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"
)
return False return False
else: else:
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip() if not os.getenv("DASHSCOPE_API_KEY", "").strip():
oa_key = os.getenv("OPENAI_API_KEY", "").strip() logger.warning(
oa_key_ok = oa_key and not ( "未设置 MODELSCOPE_API_KEY、"
oa_key.startswith("http://") or oa_key.startswith("https://") "OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
) )
if ms_key: return False
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY base = (
elif oa_key_ok: os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip(): ).strip()
logger.warning("USE_LOCAL_EMBEDDING=false 但未设置 OPENAI_EMBEDDING_MODEL") if not base:
return False logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
else: return False
if not os.getenv("DASHSCOPE_API_KEY", "").strip(): if not os.getenv("DASHSCOPE_MODEL", "").strip():
logger.warning( logger.warning("未设置 DASHSCOPE_MODEL")
"USE_LOCAL_EMBEDDING=false 但未设置 MODELSCOPE_API_KEY、" return False
"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文件(支持相对路径和绝对路径)
schema_path_str = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json") 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( orchestrator = Text2SQLOrchestrator(
schema_manager=schema_mgr, schema_manager=schema_mgr,
deepseek_config=config, deepseek_config=config,
embedding_model_path=args.embedding_model,
vector_db_path=args.vector_db, vector_db_path=args.vector_db,
max_retry=args.max_retry, max_retry=args.max_retry,
use_vector_search=not args.no_vector_search, use_vector_search=not args.no_vector_search,
@@ -366,11 +341,6 @@ def main():
help="最大重试次数(默认: 2)" help="最大重试次数(默认: 2)"
) )
_load_project_env() _load_project_env()
parser.add_argument(
"--embedding-model",
default=_embedding_model_path(),
help="Embedding模型路径(默认来自环境变量 EMBEDDING_MODEL_PATH)"
)
parser.add_argument( parser.add_argument(
"--vector-db", "--vector-db",
default=os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma").strip(), 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(默认), Embedding 封装:通过 OpenAI 兼容 ``/v1/embeddings`` 远程 API 获取向量
或可选本地 HuggingFace 目录(USE_LOCAL_EMBEDDING=true + EMBEDDING_MODEL_PATH)。 (ModelScope / OpenAI / DashScope 等由环境变量选择)。
""" """
import gc
import os 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 dataclasses import dataclass
from typing import Any, List, Optional, Union
import numpy as np
import logging
logger = logging.getLogger(__name__) 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: 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__( def __init__(
self, self,
@@ -416,7 +170,7 @@ class OpenAICompatibleRemoteEmbedding:
max_length: int = 8192, max_length: int = 8192,
show_progress: bool = False, show_progress: bool = False,
) -> np.ndarray: ) -> np.ndarray:
del max_length # API 侧截断,此处仅保持签名与本地实现一致 del max_length # API 侧截断,此处仅保持签名与历史调用方一致
if isinstance(texts, str): if isinstance(texts, str):
texts = [texts] texts = [texts]
@@ -486,42 +240,27 @@ class OpenAICompatibleRemoteEmbedding:
# 向后兼容旧名称 # 向后兼容旧名称
DashScopeOpenAIEmbedding = OpenAICompatibleRemoteEmbedding DashScopeOpenAIEmbedding = OpenAICompatibleRemoteEmbedding
# 全局单例(避免重复加载模型,节省显存/内存)
_embedding_instance: Optional[Any] = None _embedding_instance: Optional[Any] = None
def get_embedder( def get_embedder(
model_path: Optional[str] = None, _model_path: Optional[str] = None,
device: Optional[str] = None, _device: Optional[str] = None,
force_reload: bool = False, force_reload: bool = False,
) -> Any: ) -> OpenAICompatibleRemoteEmbedding:
""" """
获取 Embedding 单例:默认 ``USE_LOCAL_EMBEDDING=false``,使用远程 获取远程 Embedding 单例。``_model_path`` / ``_device`` 已废弃,仅为兼容旧调用保留。
OpenAI 兼容 / ModelScope / DashScope;为 true 时用本地目录(EMBEDDING_MODEL_PATH)。
""" """
global _embedding_instance global _embedding_instance
if force_reload or _embedding_instance is None: if force_reload or _embedding_instance is None:
if _env_flag("USE_LOCAL_EMBEDDING", "false"): _embedding_instance = OpenAICompatibleRemoteEmbedding()
_embedding_instance = Qwen3Embedding(
model_path=model_path,
device=device,
)
else:
_embedding_instance = OpenAICompatibleRemoteEmbedding()
return _embedding_instance return _embedding_instance # type: ignore[return-value]
def clear_embedder(): def clear_embedder() -> None:
"""清空单例(用于测试或切换模型)""" """清空单例(用于测试或切换远端配置)。"""
global _embedding_instance global _embedding_instance
_embedding_instance = None _embedding_instance = None
import gc
gc.collect() 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示例选择器 - 基于经验数据集动态选择相关示例 Few-shot示例选择器 - 基于经验数据集动态选择相关示例
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(本地 Qwen3 或远程 OpenAI 兼容 API, 与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(远程 OpenAI 兼容 API;
由 USE_LOCAL_EMBEDDING 及 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。 由 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
用法: 用法:
from utils.fewshot_selector import FewShotSelector from utils.fewshot_selector import FewShotSelector
@@ -26,10 +26,6 @@ import logging
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# 仅 USE_LOCAL_EMBEDDING=true 时通过 EMBEDDING_MODEL_PATH 使用;远程模式留空即可
_DEFAULT_LOCAL_EMBED_PATH = ""
@dataclass @dataclass
class ExperienceSample: class ExperienceSample:
"""经验数据样本""" """经验数据样本"""
@@ -84,7 +80,6 @@ class FewShotSelector:
def __init__( def __init__(
self, self,
samples_path: Optional[str] = None, samples_path: Optional[str] = None,
embedding_model_path: Optional[str] = None,
use_cache: bool = True, use_cache: bool = True,
use_chroma: Optional[bool] = None, use_chroma: Optional[bool] = None,
chroma_persist_dir: Optional[str] = None, chroma_persist_dir: Optional[str] = None,
@@ -93,8 +88,6 @@ class FewShotSelector:
Args: Args:
samples_path: 样本 JSONL 路径;**空字符串**表示不读文件(仅当 ``FEWSHOT_USE_CHROMA=true`` samples_path: 样本 JSONL 路径;**空字符串**表示不读文件(仅当 ``FEWSHOT_USE_CHROMA=true``
且 Chroma 中已有数据时可用)。构建索引仍请用 ``scripts/build_fewshot_chroma_index.py --samples``。 且 Chroma 中已有数据时可用)。构建索引仍请用 ``scripts/build_fewshot_chroma_index.py --samples``。
embedding_model_path: 本地 Embedding 模型目录;None 时用环境变量
EMBEDDING_MODEL_PATH(仅 USE_LOCAL_EMBEDDING=true 时有效)
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建) use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
use_chroma: 是否使用 Chroma 持久化向量库;None 时读环境变量 FEWSHOT_USE_CHROMA use_chroma: 是否使用 Chroma 持久化向量库;None 时读环境变量 FEWSHOT_USE_CHROMA
chroma_persist_dir: Chroma 目录;None 时用 FEWSHOT_CHROMA_PATH 或默认 chroma_fewshot 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") else os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in ("1", "true", "yes")
) )
self._chroma_persist_dir = chroma_persist_dir 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: if self.use_chroma:
self._init_chroma_mode() self._init_chroma_mode()
@@ -130,7 +118,7 @@ class FewShotSelector:
from utils.embedding import get_embedder from utils.embedding import get_embedder
from utils.fewshot_chroma_store import FewShotChromaStore from utils.fewshot_chroma_store import FewShotChromaStore
self._embedder = get_embedder(self._embedding_model_path) self._embedder = get_embedder()
self._chroma_store = FewShotChromaStore( self._chroma_store = FewShotChromaStore(
self._embedder, self._embedder,
persist_dir=self._chroma_persist_dir, persist_dir=self._chroma_persist_dir,
@@ -199,7 +187,7 @@ class FewShotSelector:
"""非 Chroma:初始化 Embedder 与内存 numpy 索引。""" """非 Chroma:初始化 Embedder 与内存 numpy 索引。"""
from utils.embedding import get_embedder from utils.embedding import get_embedder
self._embedder = get_embedder(self._embedding_model_path) self._embedder = get_embedder()
self._build_numpy_index() self._build_numpy_index()
def _build_numpy_index(self) -> None: def _build_numpy_index(self) -> None:
Binary file not shown.
-5
View File
@@ -94,11 +94,6 @@ hiddenimports = [
# 向量数据库 - 使用 collect_submodules 自动收集所有子模块 # 向量数据库 - 使用 collect_submodules 自动收集所有子模块
# Embedding 相关
'sentencepiece',
'accelerate',
'safetensors',
# LLM API # LLM API
'openai', 'openai',
-7
View File
@@ -34,13 +34,6 @@ dependencies = [
# 向量数据库 # 向量数据库
"chromadb>=0.4.0", "chromadb>=0.4.0",
# Embedding 模型
"transformers>=4.36.0",
"torch>=2.0.0",
"sentencepiece>=0.1.99",
"accelerate>=0.20.0",
"safetensors>=0.4.0",
# SQL 处理 # SQL 处理
"sqlglot>=20.0.0", "sqlglot>=20.0.0",
"sqlparse>=0.4.0", "sqlparse>=0.4.0",
-3
View File
@@ -9,9 +9,6 @@ openai>=1.0.0
# 向量数据库 # 向量数据库
chromadb>=0.4.0 chromadb>=0.4.0
sentencepiece>=0.1.99
accelerate>=0.20.0
safetensors>=0.4.0
# SQL 处理 # SQL 处理
sqlglot>=20.0.0 sqlglot>=20.0.0
+2 -3
View File
@@ -10,7 +10,7 @@
FEWSHOT_USE_CHROMA=true FEWSHOT_USE_CHROMA=true
FEWSHOT_CHROMA_PATH=./data/embeddings/chroma_fewshot FEWSHOT_CHROMA_PATH=./data/embeddings/chroma_fewshot
Embedding 与 Schema 向量一致,由 USE_LOCAL_EMBEDDING / OPENAI_* / MODELSCOPE_* 等决定。 Embedding 与 Schema 向量一致,由 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等决定。
""" """
from __future__ import annotations from __future__ import annotations
@@ -84,8 +84,7 @@ def main() -> int:
logger.error("未解析到任何样本") logger.error("未解析到任何样本")
return 1 return 1
embed_path = os.getenv("EMBEDDING_MODEL_PATH", "").strip() or None embedder = get_embedder()
embedder = get_embedder(embed_path)
persist = Path(args.persist_dir) persist = Path(args.persist_dir)
if not persist.is_absolute(): if not persist.is_absolute():
persist = _REPO_ROOT / persist persist = _REPO_ROOT / persist