Enhance dialog context handling for Text2SQL queries by integrating session history. Introduce new methods for summarizing previous assistant messages and determining if the last interaction was a data query. Update environment configuration for embedding options and improve error handling in SQL generation. Add user-facing delivery messages for successful SQL execution. This update supports more coherent follow-up questions and improves user experience in conversational interactions.
This commit is contained in:
+19
-22
@@ -1,6 +1,6 @@
|
||||
"""
|
||||
Embedding 封装:本地 Qwen3-Embedding,或兼容 OpenAI /v1/embeddings 的远程 API
|
||||
(ModelScope 推理、阿里云 DashScope 等)。
|
||||
Embedding 封装:兼容 OpenAI /v1/embeddings 的远程 API(默认),
|
||||
或可选本地 HuggingFace 目录(USE_LOCAL_EMBEDDING=true + EMBEDDING_MODEL_PATH)。
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -112,10 +112,10 @@ def _remote_embedding_from_env() -> _RemoteEmbeddingEnv:
|
||||
|
||||
class Qwen3Embedding:
|
||||
"""
|
||||
Qwen3-Embedding-0.6B 向量化封装
|
||||
本地 HuggingFace 格式 Embedding 模型(Mean Pooling + L2,用于向量检索)。
|
||||
|
||||
使用 Mean Pooling 将token embeddings聚合为句子向量,
|
||||
并进行L2归一化以支持余弦相似度计算。
|
||||
仅在 ``USE_LOCAL_EMBEDDING=true`` 时使用;路径由 ``model_path`` 或环境变量
|
||||
``EMBEDDING_MODEL_PATH`` 指定,**不再内置默认目录**。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -128,7 +128,7 @@ class Qwen3Embedding:
|
||||
初始化 embedding 模型
|
||||
|
||||
Args:
|
||||
model_path: 本地模型路径,若为None则从环境变量或默认路径加载
|
||||
model_path: 本地模型目录;None 或空字符串时读 ``EMBEDDING_MODEL_PATH``
|
||||
device: 推理设备('cpu', 'cuda', 'cuda:0'等),None则自动选择
|
||||
use_fp16: 是否使用FP16混合精度(GPU可用时建议开启,速度更快)
|
||||
"""
|
||||
@@ -138,22 +138,19 @@ class Qwen3Embedding:
|
||||
"pip install transformers torch sentencepiece accelerate"
|
||||
)
|
||||
|
||||
# 确定模型路径
|
||||
if model_path is None:
|
||||
model_path = os.getenv(
|
||||
"EMBEDDING_MODEL_PATH",
|
||||
"./data/models/Qwen3-Embedding-0.6B"
|
||||
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(model_path)
|
||||
model_path = Path(resolved)
|
||||
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}"
|
||||
f"本地 Embedding 模型目录不存在:{model_path}\n"
|
||||
"请修正 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false。"
|
||||
)
|
||||
|
||||
# 确定设备
|
||||
@@ -161,7 +158,7 @@ class Qwen3Embedding:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
self.device = device
|
||||
logger.info(f"加载Qwen3-Embedding模型:{model_path},设备:{device}")
|
||||
logger.info("加载本地 Embedding 模型:%s,设备:%s", model_path, device)
|
||||
|
||||
# 加载 tokenizer:fast(Rust) 解析 tokenizer.json 需较新 tokenizers;
|
||||
# 旧版本会报 ModelWrapper / untagged enum,回退到慢速 tokenizer 可恢复。
|
||||
@@ -499,13 +496,13 @@ def get_embedder(
|
||||
force_reload: bool = False,
|
||||
) -> Any:
|
||||
"""
|
||||
获取 Embedding 单例:USE_LOCAL_EMBEDDING=true 时用本地 Qwen3,否则用远程
|
||||
OpenAI 兼容 API(优先级见 OpenAICompatibleRemoteEmbedding)。
|
||||
获取 Embedding 单例:默认 ``USE_LOCAL_EMBEDDING=false``,使用远程
|
||||
OpenAI 兼容 / ModelScope / DashScope;为 true 时用本地目录(EMBEDDING_MODEL_PATH)。
|
||||
"""
|
||||
global _embedding_instance
|
||||
|
||||
if force_reload or _embedding_instance is None:
|
||||
if _env_flag("USE_LOCAL_EMBEDDING", "true"):
|
||||
if _env_flag("USE_LOCAL_EMBEDDING", "false"):
|
||||
_embedding_instance = Qwen3Embedding(
|
||||
model_path=model_path,
|
||||
device=device,
|
||||
|
||||
Reference in New Issue
Block a user