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:
陈辅元
2026-04-14 18:02:12 +08:00
parent 026d3bf88c
commit ca8bc5e7de
87 changed files with 5808 additions and 6351 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
# Text-to-SQL Multi-Agent System
# 基于 CAMEL AI + DeepSeek + Qwen3-Embedding 的智能SQL生成系统
# 基于 CAMEL AI + DeepSeek + 向量 Embedding 的智能 SQL 生成系统
__version__ = "0.1.0"
__author__ = "Backman Team"
Binary file not shown.
+102 -19
View File
@@ -64,7 +64,7 @@ class Text2SQLOrchestrator:
schema_manager: Schema管理器实例
deepseek_api_key: DeepSeek API密钥(也可通过环境变量DEEPSEEK_API_KEY)
deepseek_config: DeepSeek配置对象(优先于api_key)
embedding_model_path: Qwen3-Embedding模型路径
embedding_model_path: 本地 Embedding 模型目录(仅 USE_LOCAL_EMBEDDING=true)
vector_db_path: 向量数据库路径
max_retry: 最大重试次数(包含首次生成)
use_vector_search: 是否使用向量检索粗筛
@@ -96,12 +96,22 @@ class Text2SQLOrchestrator:
if self.fewshot_enabled:
try:
path = fewshot_samples_path or os.getenv(
"FEWSHOT_DATA_PATH",
"./data/experiences/all_samples.jsonl"
use_chroma = os.getenv("FEWSHOT_USE_CHROMA", "false").lower() in (
"1",
"true",
"yes",
)
if fewshot_samples_path is not None:
path = str(fewshot_samples_path).strip()
else:
env_p = os.getenv("FEWSHOT_DATA_PATH")
if env_p is not None:
path = env_p.strip()
else:
# Chroma 优先时默认不再依赖 JSONL;否则保留原默认路径
path = "" if use_chroma else "./data/experiences/all_samples.jsonl"
self.fewshot_selector = FewShotSelector(
path,
path or None,
embedding_model_path=self._embedding_model_path,
)
logger.info(
@@ -119,6 +129,23 @@ class Text2SQLOrchestrator:
+ (", en→zh=on" if self.translate_english_to_zh else ", en→zh=off")
)
@staticmethod
def _merge_dialog_for_model(
dialog_context: str, question: str, max_len: int
) -> str:
"""拼接上文与当前问句,控制总长,优先保留当前问句完整。"""
dc = (dialog_context or "").strip()
q = (question or "").strip()
if not dc:
return q
tail = "\n\n【当前用户问题】\n" + q
if len(dc) + len(tail) <= max_len:
return dc + tail
room = max_len - len(tail)
if room < 80:
return tail[-max_len:]
return dc[:room].rstrip() + tail
def _get_vector_index(self) -> SchemaIndexer:
"""获取或创建向量索引(懒加载)"""
if self._vector_index is None:
@@ -291,6 +318,7 @@ class Text2SQLOrchestrator:
schema_str: str,
dialect: str = "tsql",
validation_feedback: Optional[str] = None,
dialog_context: Optional[str] = None,
) -> str:
"""
SQL生成(SQL Generator Agent)
@@ -301,17 +329,24 @@ class Text2SQLOrchestrator:
dialect: SQL方言
validation_feedback: 非空时附加到用户提示(重试时传入上次校验错误)
dialog_context: 前几轮对话摘要;与 ``question`` 一并供指代消解与续问。
Returns:
SQL语句
"""
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
from utils.sql_parser import normalize_sql_for_dialect
dc = (dialog_context or "").strip()
fewshot_question = question
if dc:
fewshot_question = f"{dc}\n\n【当前问】{question}"
# Few-shot 增强
if self.fewshot_enabled and self.fewshot_selector:
try:
examples = self.fewshot_selector.select(
question=question,
question=fewshot_question,
top_k=self.fewshot_top_k,
min_rating=self.fewshot_min_rating
)
@@ -329,7 +364,15 @@ class Text2SQLOrchestrator:
if dialect == "tsql":
dialect_label = "Microsoft SQL Server (T-SQL)"
user_content = SQL_GENERATOR_USER.format(
prefix = ""
if dc:
prefix = (
"【对话上文】(用于理解「这/那/同样/上面/刚才」等指代及续问条件;"
"请结合下文「当前用户问题」生成 SQL。)\n"
f"{dc}\n\n"
)
user_content = prefix + SQL_GENERATOR_USER.format(
schema=schema_str,
question=question,
dialect=dialect_label,
@@ -380,6 +423,7 @@ class Text2SQLOrchestrator:
schema_str: str,
dialect: str = "tsql",
question: str = "",
dialog_context: Optional[str] = None,
) -> Tuple[bool, List[str], List[str], Optional[int], Optional[str]]:
"""
SQL验证(程序验证 + 库探针 + 按探针分支的 LLM)
@@ -389,6 +433,7 @@ class Text2SQLOrchestrator:
schema_str: Schema描述
dialect: 与生成一致的 SQL 方言(sqlglot 名,默认 tsql)
question: 用户自然语言(探针为 0 时用于生成补充说明)
dialog_context: 会话上文;探针 0 时与 question 一并传入说明模型
Returns:
(是否通过, 错误列表, 警告列表, 库执行探针状态, 无数据时的用户说明)
@@ -433,11 +478,11 @@ class Text2SQLOrchestrator:
db_execution_status, db_probe_err = probe_sql_execution_status_ex(sql)
if db_execution_status == -1:
msg = (
"数据库执行验证失败:SQL 在目标库执行报错(探针状态 -1),"
"将据此重新生成 SQL。"
"【库探针结果:-1 执行失败】SQL 在目标库执行报错,"
"系统将依据下列错误**自动重新生成** SQL(请等待重试结果)。"
)
if db_probe_err:
msg += f" 数据库返回:{db_probe_err}"
msg += f"\n数据库返回:{db_probe_err}"
errors.append(msg)
elif db_execution_status is None:
warnings.append(
@@ -482,23 +527,35 @@ class Text2SQLOrchestrator:
# === 阶段2b:探针 0 时生成用户可读补充说明(仍返回 SQL,由 API/CLI 一并展示) ===
if db_execution_status == 0 and len(errors) == 0:
prefix = (
"该 SQL 已在数据库成功执行,但返回的数据行数为 0(未查到匹配记录)。"
"请将下方 SQL 与说明一并核对;若不符合预期,请补充或调整条件后再次提问。"
"【库探针结果:0 行】该 SQL 已在数据库成功执行,但**返回数据行数为 0**(未查到匹配记录)。"
"下方已附带完整 SQL 与原因分析,请一并阅读。"
)
fb_q = question
dc = (dialog_context or "").strip()
if dc:
fb_q = f"{dc}\n\n【当前用户问题】\n{question}"
try:
llm_fb = self.deepseek.empty_result_user_feedback(
question=question,
question=fb_q,
sql=sql,
schema=schema_str,
)
empty_feedback = f"{prefix}\n\n【分析与建议】\n{llm_fb}"
empty_feedback = f"{prefix}\n\n【问题分析】\n{llm_fb}"
except Exception as e:
logger.warning(f"无数据说明生成失败: {e}")
empty_feedback = (
f"{prefix}\n\n【分析与建议】\n"
f"{prefix}\n\n【问题分析】\n"
"未能自动生成详细分析。请补充时间范围、筛选条件或业务对象后重新提问。"
)
follow = (
"\n\n【追问 — 请补充后再次提问以重新生成 SQL】\n"
"1. 请根据上述分析,尽量具体地补充或修正:**时间范围**、**业务对象**(账户/合约/代码等)、"
"**筛选口径** 或 **您认为 SQL 中不合理的条件**。\n"
"2. 补充说明后请**重新发起一次自然语言提问**(无需粘贴 SQL),系统会结合您的新描述**重新生成**查询。"
)
empty_feedback = (empty_feedback or prefix) + follow
is_valid = len(errors) == 0
return is_valid, errors, warnings, db_execution_status, empty_feedback
@@ -507,7 +564,8 @@ class Text2SQLOrchestrator:
question: str,
dialect: str = "tsql",
top_k_candidates: int = 20,
include_schema_in_result: bool = False
include_schema_in_result: bool = False,
dialog_context: Optional[str] = None,
) -> GenerationResult:
"""
主生成流程
@@ -517,6 +575,7 @@ class Text2SQLOrchestrator:
dialect: SQL方言(默认 tsql;与 sqlglot 一致)
top_k_candidates: 粗筛候选表数量
include_schema_in_result: 结果中是否包含使用的Schema字符串
dialog_context: 前几轮对话可读摘要;选表、向量粗筛、SQL 生成与无数据说明会参考
Returns:
GenerationResult对象
@@ -543,6 +602,9 @@ class Text2SQLOrchestrator:
logger.warning("[GEN] 英译中失败,使用原文: %s", e)
question = work_question
dc_raw = (dialog_context or "").strip()
retrieval_question = self._merge_dialog_for_model(dc_raw, question, 4000)
linker_question = self._merge_dialog_for_model(dc_raw, question, 6000)
logger.info(f"[GEN] 开始生成SQL:{question[:50]}...")
attempt = 0
@@ -558,14 +620,18 @@ class Text2SQLOrchestrator:
# === Step 1: Schema筛选(仅首次) ===
if attempt == 0:
# 1.1 粗筛
candidate_tables = self._coarse_filter(question, top_k=top_k_candidates)
candidate_tables = self._coarse_filter(
retrieval_question, top_k=top_k_candidates
)
# 1.2 LLM精筛
relevant_tables, reasoning = self._llm_select_tables(
question,
linker_question,
candidate_tables,
)
relevant_tables = self._prioritize_broker_tables(question, relevant_tables)
relevant_tables = self._prioritize_broker_tables(
linker_question, relevant_tables
)
# 1.3 外键扩展
expanded_tables = self._expand_relations(relevant_tables)
@@ -613,6 +679,7 @@ class Text2SQLOrchestrator:
filtered_schema_str,
dialect,
validation_feedback=feedback,
dialog_context=dc_raw or None,
)
last_sql = sql
except Exception as e:
@@ -626,6 +693,7 @@ class Text2SQLOrchestrator:
filtered_schema_str,
dialect=dialect,
question=question,
dialog_context=dc_raw or None,
)
if db_probe is not None:
last_db_execution_status = db_probe
@@ -643,6 +711,19 @@ class Text2SQLOrchestrator:
meta["db_execution_status"] = db_probe
if empty_feedback:
meta["db_empty_feedback"] = empty_feedback
if db_probe == 1:
try:
meta["sql_delivery_message"] = (
self.deepseek.sql_probe_success_delivery_message(
question=question,
sql=sql,
)
)
except Exception as e:
logger.warning("[GEN] 探针1交付说明生成失败: %s", e)
meta["sql_delivery_message"] = None
if dc_raw:
meta["dialog_context_chars"] = len(dc_raw)
result = GenerationResult(
sql=sql,
@@ -663,6 +744,8 @@ class Text2SQLOrchestrator:
fail_meta: Dict = dict(translation_meta)
if last_db_execution_status is not None:
fail_meta["db_execution_status"] = last_db_execution_status
if dc_raw:
fail_meta["dialog_context_chars"] = len(dc_raw)
return GenerationResult(
sql=last_sql or "",
valid=False,
Binary file not shown.
Binary file not shown.
+43
View File
@@ -283,6 +283,24 @@ EMPTY_RESULT_FEEDBACK_USER = """用户原始问题:
请直接输出给终端用户阅读的说明文字(纯文本)。"""
# 库探针为 1(有返回行)时,面向用户的 SQL 交付说明(与下方 SQL 一并展示)
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM = """你是证券/期货类业务库的 Text2SQL 助手。
用户的自然语言问题已转为 SQL,且在目标库**试执行成功且至少返回一行数据**。
请用 1~3 句简洁中文向用户说明:
- 该 SQL 大致在查询或统计什么(业务语义);
- 可提示用户可在下方查看完整 SQL 并自行执行或导出。
不要编造具体数据值;不要逐列复述;不要输出 Markdown 代码块或 JSON。"""
SQL_PROBE_SUCCESS_DELIVERY_USER = """用户问题:
{question}
已通过库探针(有数据行)的 SQL:
{sql}
请输出面向终端用户的简短说明(纯文本)。"""
# ========== Few-Shot 示例 ==========
FEW_SHOT_EXAMPLES: Dict[str, str] = {
@@ -350,3 +368,28 @@ TRANSLATE_NL_TO_ZH_USER = """原句:
{question}
仅输出一句中文:"""
# ========== 对话意图分类(LLM,与规则分类配合)==========
DIALOG_INTENT_CLASSIFIER_SYSTEM = """你是证券/期货类 Text2SQL 产品的「意图分类」模块。
只输出**一行**严格 JSON(不要 markdown、不要解释),格式:
{"intent":"text2sql"|"conversation","reply_zh":"..."}
字段含义:
- intent=text2sql:用户本轮在**提出新的或可执行的数据查询/统计需求**(含对上一轮的**具体补充条件**,如「再加上经纪商维度」「改成按日」),需要走 SQL 生成。
- intent=conversation:用户在做**寒暄致谢**、**元问题**(你是谁)、**对结果的情绪反馈但缺少可执行信息**(如「不对呀」「结果错了」「重新算」却未说清要改什么)、**纯抱怨或否定而无新条件**——此时不应生成 SQL;reply_zh 用简短中文引导用户**具体说明**要查什么或错在哪里。
- reply_zh:当 intent=conversation 时必填,为直接展示给用户的友好中文(1~4 句);intent=text2sql 时填空字符串 ""。
注意:若「上一轮助手刚返回过数据/SQL」而用户只说结果不对、未给出新的筛选/维度/时间,判为 conversation。"""
DIALOG_INTENT_CLASSIFIER_USER = """上一轮助手是否为「数据查询/SQL 结果」:{last_turn_label}
可选会话摘要(可能为空):
---
{context_snip}
---
用户本轮输入:
{user_message}
只输出 JSON:"""
+1 -1
View File
@@ -17,7 +17,7 @@ class Settings(BaseSettings):
max_tokens: int = 4096
# Embedding配置
embedding_model_path: str = "./data/models/Qwen3-Embedding-0.6B"
embedding_model_path: str = "" # 仅本地 Embedding 时使用;默认走远程 API
vector_dim: int = 2048
use_local_embedding: bool = True
+37
View File
@@ -253,6 +253,43 @@ class DeepSeekClient:
text = (msg.content or "").strip()
return text
def sql_probe_success_delivery_message(
self,
question: str,
sql: str,
*,
max_sql_chars: int = 4000,
**kwargs: Any,
) -> str:
"""
库探针为 1(执行成功且至少一行数据)时,生成面向用户的简短交付说明。
"""
from config.prompts import (
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM,
SQL_PROBE_SUCCESS_DELIVERY_USER,
)
q = (question or "").strip() or "(无)"
s = (sql or "").strip()
if len(s) > max_sql_chars:
s = s[: max_sql_chars - 20].rstrip() + "\n-- …(已截断)"
messages = [
{"role": "system", "content": SQL_PROBE_SUCCESS_DELIVERY_SYSTEM},
{
"role": "user",
"content": SQL_PROBE_SUCCESS_DELIVERY_USER.format(
question=q,
sql=s,
),
},
]
kwargs.setdefault("temperature", 0.2)
kwargs.setdefault("top_p", 1.0)
kwargs.setdefault("max_tokens", 320)
msg = self.chat(messages, **kwargs)
return (msg.content or "").strip()
def select_tables(
self,
question: str,
+16 -15
View File
@@ -27,9 +27,6 @@ logging.basicConfig(
)
logger = logging.getLogger(__name__)
_DEFAULT_EMBEDDING_PATH = "./data/models/Qwen3-Embedding-0.6B"
def _repo_root() -> Path:
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
return Path(__file__).resolve().parent.parent
@@ -41,7 +38,8 @@ def _load_project_env():
def _embedding_model_path() -> str:
return os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_EMBEDDING_PATH).strip()
"""仅 ``USE_LOCAL_EMBEDDING=true`` 时需要;远程 Embedding 可为空。"""
return os.getenv("EMBEDDING_MODEL_PATH", "").strip()
def resolve_sql_dialect(name: str) -> str:
@@ -55,20 +53,23 @@ def resolve_sql_dialect(name: str) -> str:
def setup_environment():
"""环境检查"""
_load_project_env()
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "true").strip().lower() in (
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in (
"1", "true", "yes", "on",
)
if use_local_emb:
# 与 .env 中 EMBEDDING_MODEL_PATH 及 utils.embedding 一致
model_path = Path(_embedding_model_path())
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(f"Embedding模型不存在: {model_path}")
logger.info("请先下载模型:")
logger.info(" modelscope download --model 'Qwen/Qwen3-Embedding-0.6B' "
f"--local_dir '{model_path}'")
logger.info("或使用远程 Embedding API:USE_LOCAL_EMBEDDING=false,并配置 "
"MODELSCOPE_API_KEY、或 OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL"
"(及可选 OPENAI_BASE_URL)、或 DASHSCOPE_*(百炼)")
logger.warning("本地 Embedding 目录不存在: %s", model_path)
logger.info(
"请设置正确的 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false,"
"并配置 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding"
)
return False
else:
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
@@ -206,7 +207,7 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
logger.info(f"[Q] 问题: {question}")
classified = classify_dialog(question)
classified = classify_dialog(question, llm_client=orchestrator.deepseek)
if classified.intent == DialogIntent.CONVERSATION:
reply = classified.reply_suggestion or ""
logger.info("[Q] 意图: conversation(跳过 SQL 生成)")
Binary file not shown.
+31 -8
View File
@@ -40,6 +40,7 @@ class SchemaIndexer:
self.embedder = embedder
self.persist_dir = Path(persist_dir)
self.persist_dir.mkdir(parents=True, exist_ok=True)
self.collection_name = collection_name
# 初始化ChromaDB客户端
self.client = chromadb.PersistentClient(
@@ -49,11 +50,11 @@ class SchemaIndexer:
# 获取或创建集合
self.collection = self.client.get_or_create_collection(
name=collection_name,
name=self.collection_name,
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
)
logger.info(f"[OK] 初始化SchemaIndexer: collection={collection_name}")
logger.info(f"[OK] 初始化SchemaIndexer: collection={self.collection_name}")
def build_index(
self,
@@ -80,7 +81,12 @@ class SchemaIndexer:
if force_rebuild and existing_ids:
logger.info(f"强制重建索引,删除{len(existing_ids)}条旧记录")
self.collection.delete()
# Chroma 新版本要求 delete 必须带 ids/where;整库清空用删集合再建
self.client.delete_collection(self.collection_name)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
# 准备数据
table_texts = []
@@ -97,7 +103,9 @@ class SchemaIndexer:
# 批量计算embedding
logger.info(f"计算{len(table_texts)}张表的embedding...")
embeddings = self.embedder.encode(table_texts, batch_size=batch_size)
embeddings = self.embedder.encode(
table_texts, batch_size=batch_size, normalize=True
)
# 存入ChromaDB
self.collection.add(
@@ -136,8 +144,8 @@ class SchemaIndexer:
...
]
"""
# 编码查询文本
query_embedding = self.embedder.encode([query])
# 编码查询文本(与建库时一致:L2 归一化 + 余弦空间)
query_embedding = self.embedder.encode([query], normalize=True)
# 执行检索
results = self.collection.query(
@@ -205,7 +213,12 @@ class SchemaIndexer:
def clear(self):
"""清空索引"""
self.collection.delete()
if self.collection.count() > 0:
self.client.delete_collection(self.collection_name)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
logger.info("索引已清空")
def count(self) -> int:
@@ -218,8 +231,18 @@ class SchemaIndexer:
"""获取索引统计信息"""
count = self.collection.count()
result = self.collection.get(include=["metadatas"])
metas = result.get("metadatas") or []
total_columns = sum(m.get("column_count", 0) for m in result["metadatas"])
def _col_count(m: Optional[Dict]) -> int:
if not m:
return 0
v = m.get("column_count", 0)
try:
return int(v)
except (TypeError, ValueError):
return 0
total_columns = sum(_col_count(m) for m in metas)
return {
"indexed_tables": count,
Binary file not shown.
+158 -6
View File
@@ -1,17 +1,27 @@
"""
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
支持环境变量 ``DIALOG_INTENT_CLASSIFIER``:
- ``rules``:仅用关键词与短语规则(无 LLM 调用)。
- ``hybrid``(默认):明显查数词/寒暄/元问题走规则;其余交 LLM 判断(需调用方传入 ``llm_client``)。
"""
from __future__ import annotations
import logging
import os
import re
import unicodedata
from enum import Enum
from typing import NamedTuple, Optional
from typing import TYPE_CHECKING, NamedTuple, Optional
if TYPE_CHECKING:
from llm.deepseek_client import DeepSeekClient
logger = logging.getLogger(__name__)
_INTENT_CLASSIFIER_ENV = "DIALOG_INTENT_CLASSIFIER"
class DialogIntent(str, Enum):
TEXT2SQL = "text2sql"
@@ -100,12 +110,24 @@ def _strip_trailing_punct(t: str) -> str:
return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip()
def classify_dialog(user_text: str) -> DialogClassifyResult:
"""
对用户一轮输入做粗分类。
def _intent_classifier_mode() -> str:
m = (os.getenv(_INTENT_CLASSIFIER_ENV) or "hybrid").strip().lower()
if m not in ("rules", "hybrid"):
logger.warning(
"[dialog] unknown DIALOG_INTENT_CLASSIFIER=%r, use hybrid", m
)
return "hybrid"
return m
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 CONVERSATION;
其余默认 TEXT2SQL,避免漏判真实查询。
def _classify_dialog_rules(
user_text: str,
*,
last_turn_was_data_query: bool = False,
) -> DialogClassifyResult:
"""
规则分类(与历史行为一致):优先查数关键词;寒暄/元问题为 conversation;
若上一轮为数据查询且非寒暄/元问题,则倾向 text2sql。
"""
t = _normalize(user_text)
if not t:
@@ -131,5 +153,135 @@ def classify_dialog(user_text: str) -> DialogClassifyResult:
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
)
if last_turn_was_data_query:
logger.debug("[dialog] intent=text2sql (follow-up after data query)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
logger.debug("[dialog] intent=text2sql (default)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
def _classify_dialog_llm(
client: DeepSeekClient,
user_text: str,
*,
last_turn_was_data_query: bool,
dialog_context: Optional[str],
) -> DialogClassifyResult:
from config.prompts import (
DIALOG_INTENT_CLASSIFIER_SYSTEM,
DIALOG_INTENT_CLASSIFIER_USER,
)
last_label = "是" if last_turn_was_data_query else "否"
ctx = (dialog_context or "").strip()
if len(ctx) > 2400:
ctx = ctx[:2400].rstrip() + "\n…(已截断)"
if not ctx:
ctx = "(无)"
raw = client.chat_with_json(
[
{"role": "system", "content": DIALOG_INTENT_CLASSIFIER_SYSTEM},
{
"role": "user",
"content": DIALOG_INTENT_CLASSIFIER_USER.format(
last_turn_label=last_label,
context_snip=ctx,
user_message=_normalize(user_text),
),
},
],
temperature=0.0,
max_tokens=256,
top_p=1.0,
)
if raw.get("_json_decode_failed"):
raise ValueError("intent JSON decode failed")
intent_s = (raw.get("intent") or "").strip().lower()
reply = (raw.get("reply_zh") or "").strip()
if intent_s == "conversation":
if not reply:
reply = (
"若上一版 SQL 或结果不符合预期,请具体说明:希望增加/修改哪些条件、"
"时间范围或统计维度,以便重新生成。"
)
logger.info("[dialog] intent=conversation (LLM)")
return DialogClassifyResult(DialogIntent.CONVERSATION, reply_suggestion=reply)
if intent_s == "text2sql":
logger.info("[dialog] intent=text2sql (LLM)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
raise ValueError(f"unexpected intent field: {intent_s!r}")
def classify_dialog(
user_text: str,
*,
last_turn_was_data_query: bool = False,
dialog_context: Optional[str] = None,
llm_client: Optional[DeepSeekClient] = None,
) -> DialogClassifyResult:
"""
对用户一轮输入做分类。
``DIALOG_INTENT_CLASSIFIER``:
- ``rules``:仅规则。
- ``hybrid``:规则快速命中查数词/寒暄/元问题后返回;否则在有 ``llm_client`` 时用 LLM,
失败或无客户端时回退规则。
Args:
user_text: 用户输入。
last_turn_was_data_query: 上一轮助手是否为数据查询(供规则与 LLM 参考)。
dialog_context: 会话摘要,供 LLM 参考(可选)。
llm_client: DeepSeek 客户端;hybrid 下 LLM 分支需要。
"""
mode = _intent_classifier_mode()
if mode == "rules":
return _classify_dialog_rules(
user_text, last_turn_was_data_query=last_turn_was_data_query
)
# hybrid
t = _normalize(user_text)
if not t:
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=_EMPTY_INPUT_REPLY
)
if _SQL_OR_QUERY_HINT_RE.search(t):
logger.debug("[dialog] intent=text2sql (query/business hint, hybrid fast)")
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
core = _strip_trailing_punct(t)
if core.casefold() in _CHITCHAT_KEYS:
logger.debug("[dialog] intent=conversation (chitchat phrase, hybrid fast)")
return DialogClassifyResult(
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
)
if _META_QUESTION_RE.search(t):
logger.debug("[dialog] intent=conversation (meta question, hybrid fast)")
return DialogClassifyResult(
DialogIntent.CONVERSATION,
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
)
if llm_client is not None:
try:
return _classify_dialog_llm(
llm_client,
user_text,
last_turn_was_data_query=last_turn_was_data_query,
dialog_context=dialog_context,
)
except Exception as e:
logger.warning("[dialog] LLM intent failed, fallback rules: %s", e)
return _classify_dialog_rules(
user_text, last_turn_was_data_query=last_turn_was_data_query
)
+130
View File
@@ -0,0 +1,130 @@
"""
从 Lite NL 会话消息构造 Text2SQL 可用的上文摘要,并判断上一轮是否为数据查询(用于续问意图分类)。
"""
from __future__ import annotations
import json
import logging
from typing import Any, Dict, List, Tuple
logger = logging.getLogger(__name__)
_MAX_ASSISTANT_SQL_CHARS = 1400
_MAX_ASSISTANT_TEXT_CHARS = 900
_MAX_BLOCK_CHARS = 7500
def _truncate(s: str, max_len: int) -> str:
s = (s or "").strip()
if len(s) <= max_len:
return s
return s[: max_len - 1].rstrip() + "…"
def summarize_assistant_nl_payload(content: str) -> str:
"""将落库的 assistant JSON 转为一小段可读摘要。"""
raw = (content or "").strip()
if not raw:
return ""
try:
data: Dict[str, Any] = json.loads(raw)
except json.JSONDecodeError:
return _truncate(raw, _MAX_ASSISTANT_TEXT_CHARS)
br = data.get("branch_result") or {}
if isinstance(br.get("answer"), str) and br["answer"].strip():
return "[助手] " + _truncate(br["answer"].strip(), _MAX_ASSISTANT_TEXT_CHARS)
sql = (br.get("sql") or "").strip() if isinstance(br.get("sql"), str) else ""
if sql:
parts = ["[上轮 SQL] " + _truncate(sql, _MAX_ASSISTANT_SQL_CHARS)]
if br.get("follow_up_required"):
parts.append("[状态] 上轮为库探针0行,需补充条件后重问")
for key in ("db_empty_feedback", "sql_explain"):
v = br.get(key)
if isinstance(v, str) and v.strip():
parts.append("[说明] " + _truncate(v.strip(), 500))
break
return "\n".join(parts)
intent = (data.get("intent") or {}).get("intent")
return _truncate(f"[助手] intent={intent}", _MAX_ASSISTANT_TEXT_CHARS)
def last_assistant_was_data_query(items: List[Dict[str, Any]]) -> bool:
"""最后一条 assistant 消息是否为数据查询分支(含 SQL 或 DATA_QUERY intent)。"""
for m in reversed(items or []):
if m.get("role") != "assistant":
continue
raw = (m.get("content") or "").strip()
if not raw:
return False
try:
data = json.loads(raw)
except json.JSONDecodeError:
return False
intent = (data.get("intent") or {}).get("intent")
if intent == "DATA_QUERY":
return True
br = data.get("branch_result") or {}
sql = (br.get("sql") or "").strip() if isinstance(br.get("sql"), str) else ""
return bool(sql)
return False
def messages_to_text2sql_context(
items: List[Dict[str, Any]],
*,
max_chars: int = _MAX_BLOCK_CHARS,
max_pairs: int = 8,
) -> Tuple[str, int]:
"""
将会话消息列表转为供模型阅读的「上文」文本(不含本轮用户输入)。
Returns:
(context_text, num_user_turns_included)
"""
if not items:
return "", 0
pairs: List[Tuple[str, str]] = []
i = 0
n = len(items)
while i < n:
u = items[i]
if u.get("role") != "user":
i += 1
continue
user_text = (u.get("content") or "").strip()
asst_text = ""
if i + 1 < n and items[i + 1].get("role") == "assistant":
asst_text = summarize_assistant_nl_payload(str(items[i + 1].get("content") or ""))
i += 2
else:
i += 1
if user_text or asst_text:
pairs.append((user_text, asst_text))
if not pairs:
return "", 0
pairs = pairs[-max_pairs:]
lines: List[str] = []
total = 0
user_count = 0
for idx, (uq, aq) in enumerate(pairs, start=1):
chunk_parts = [f"--- 第{idx}轮 ---"]
if uq:
chunk_parts.append(f"用户:{uq}")
if aq:
chunk_parts.append(f"{aq}")
chunk = "\n".join(chunk_parts)
if total + len(chunk) + 2 > max_chars:
break
lines.append(chunk)
total += len(chunk) + 2
if uq:
user_count += 1
return "\n\n".join(lines).strip(), user_count
+19 -22
View File
@@ -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,
+214
View File
@@ -0,0 +1,214 @@
"""
Few-shot 经验样本的 Chroma 向量库存储与检索。
与 SchemaIndexer 一致:使用项目统一 ``get_embedder()``,持久化目录默认
``./data/embeddings/chroma_fewshot``,集合名 ``fewshot_samples``。
"""
from __future__ import annotations
import json
import logging
import os
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import chromadb
from chromadb.config import Settings as ChromaSettings
logger = logging.getLogger(__name__)
DEFAULT_FEWSHOT_CHROMA_DIR = "./data/embeddings/chroma_fewshot"
COLLECTION_NAME = "fewshot_samples"
# Chroma metadata 单值不宜过大,SQL/说明超长时截断
_MAX_SQL_META = 16000
_MAX_EXPLAIN_META = 6000
_MAX_QEN_META = 2000
class FewShotChromaStore:
"""Few-shot JSONL → Chroma 写入与按向量检索。"""
def __init__(
self,
embedder: Any,
persist_dir: Optional[str] = None,
collection_name: Optional[str] = None,
):
self.embedder = embedder
# Chroma 的 collection 名;显式参数优先,否则读 FEWSHOT_CHROMA_COLLECTION,再回退默认
explicit = (collection_name or "").strip()
from_env = (os.getenv("FEWSHOT_CHROMA_COLLECTION") or "").strip()
self.collection_name = explicit or from_env or COLLECTION_NAME
self.persist_dir = Path(
(persist_dir or os.getenv("FEWSHOT_CHROMA_PATH") or DEFAULT_FEWSHOT_CHROMA_DIR).strip()
)
self.persist_dir.mkdir(parents=True, exist_ok=True)
self.client = chromadb.PersistentClient(
path=str(self.persist_dir),
settings=ChromaSettings(anonymized_telemetry=False),
)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
logger.info(
"[OK] FewShotChromaStore: path=%s collection=%s count=%s",
self.persist_dir,
self.collection_name,
self.collection.count(),
)
def count(self) -> int:
return int(self.collection.count())
def clear(self) -> None:
"""删除集合内全部文档(用于 force rebuild)。"""
if self.collection.count() > 0:
self.client.delete_collection(self.collection_name)
self.collection = self.client.get_or_create_collection(
name=self.collection_name,
metadata={"hnsw:space": "cosine"},
)
logger.info("[OK] Few-shot Chroma 集合已清空并重建")
def build_from_samples(
self,
samples: List["ExperienceSample"],
*,
force_rebuild: bool = False,
batch_size: int = 32,
) -> int:
"""
将已解析的样本列表写入 Chroma(向量由 question_zh 编码)。
Returns:
写入条数
"""
if not samples:
logger.warning("Few-shot Chroma: 无样本,跳过构建")
return 0
if force_rebuild and self.count() > 0:
self.clear()
if not force_rebuild and self.count() > 0:
logger.info("Few-shot Chroma 已有 %s 条,跳过构建(加 --force 可重建)", self.count())
return self.count()
texts = [(s.question_zh or "").strip() or " " for s in samples]
# Chroma 要求 ids 为唯一字符串
ids = [str(s.qid) for s in samples]
metadatas: List[Dict[str, Any]] = []
for s in samples:
metadatas.append(
{
"qid": str(s.qid),
"question_en": ((s.question_en or "")[:_MAX_QEN_META]),
"sql": (s.sql or "")[:_MAX_SQL_META],
"explanation": (s.explanation or "")[:_MAX_EXPLAIN_META],
"rating": int(s.rating) if s.rating is not None else -1,
"difficulty": (s.difficulty or "medium")[:32],
"tags_json": json.dumps(s.tags or [], ensure_ascii=False)[:4000],
}
)
logger.info("Few-shot Chroma: 计算 %s 条 embedding...", len(texts))
embeddings = self.embedder.encode(
texts,
batch_size=min(batch_size, len(texts)),
normalize=True,
show_progress=True,
)
self.collection.add(
embeddings=embeddings.tolist(),
documents=texts,
metadatas=metadatas,
ids=ids,
)
logger.info("[OK] Few-shot Chroma 索引完成: %s 条", len(ids))
return len(ids)
def search_raw(
self,
query: str,
top_k: int = 20,
) -> List[Tuple[float, Dict[str, Any], str]]:
"""
Returns:
(score, metadata_dict, document) 列表,score 同 SchemaIndexer 为 1-distance
"""
n = self.count()
if n == 0:
return []
q_emb = self.embedder.encode(
[(query or "").strip() or " "],
batch_size=1,
normalize=True,
show_progress=False,
)
k = min(max(top_k, 1), n)
results = self.collection.query(
query_embeddings=q_emb.tolist(),
n_results=k,
)
out: List[Tuple[float, Dict[str, Any], str]] = []
if not results["ids"] or not results["ids"][0]:
return out
for tid, dist, meta, doc in zip(
results["ids"][0],
results["distances"][0],
results["metadatas"][0],
results["documents"][0],
):
score = 1.0 - float(dist)
m = dict(meta or {})
m["qid"] = tid
out.append((score, m, doc or ""))
return out
def get_all_as_samples(self) -> List[Any]:
"""导出集合中全部样本(用于统计/按标签列举;条数大时慎用)。"""
if self.count() == 0:
return []
data = self.collection.get(include=["metadatas", "documents"])
ids = data.get("ids") or []
metas = data.get("metadatas") or []
docs = data.get("documents") or []
out: List[Any] = []
for i, qid in enumerate(ids):
meta = dict(metas[i] or {}) if i < len(metas) else {}
meta["qid"] = qid
doc = docs[i] if i < len(docs) else ""
out.append(sample_from_chroma_metadata(meta, doc))
return out
def sample_from_chroma_metadata(meta: Dict[str, Any], document: str) -> "ExperienceSample":
from utils.fewshot_selector import ExperienceSample
tags_raw = meta.get("tags_json") or "[]"
try:
tags = json.loads(tags_raw) if isinstance(tags_raw, str) else []
except json.JSONDecodeError:
tags = []
if not isinstance(tags, list):
tags = []
r = meta.get("rating")
try:
ri = int(r) if r is not None else -1
except (TypeError, ValueError):
ri = -1
rating = ri if ri >= 0 else None
qen = meta.get("question_en")
return ExperienceSample(
qid=str(meta.get("qid", "")),
question_zh=document or str(meta.get("question_zh", "")),
question_en=qen if qen else None,
sql=str(meta.get("sql", "")),
explanation=str(meta.get("explanation", "")),
rating=rating,
tags=tags,
difficulty=str(meta.get("difficulty", "medium") or "medium"),
)
+162 -20
View File
@@ -7,7 +7,9 @@ Few-shot示例选择器 - 基于经验数据集动态选择相关示例
用法:
from utils.fewshot_selector import FewShotSelector
selector = FewShotSelector("data/experiences/all_samples.jsonl")
# Chroma 模式(FEWSHOT_USE_CHROMA=true):可不传 JSONL,路径传 None 或 ""
selector = FewShotSelector(None)
# 或传统:FewShotSelector("data/experiences/all_samples.jsonl")
examples = selector.select(question="查询2024年1月的销售额", top_k=3)
# 在Prompt中使用
@@ -24,7 +26,8 @@ import logging
logger = logging.getLogger(__name__)
_DEFAULT_LOCAL_EMBED_PATH = "./data/models/Qwen3-Embedding-0.6B"
# 仅 USE_LOCAL_EMBEDDING=true 时通过 EMBEDDING_MODEL_PATH 使用;远程模式留空即可
_DEFAULT_LOCAL_EMBED_PATH = ""
@dataclass
@@ -80,50 +83,128 @@ class FewShotSelector:
def __init__(
self,
samples_path: str,
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,
):
"""
Args:
samples_path: 样本JSONL文件路径
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
"""
self.samples_path = Path(samples_path)
sp = (samples_path or "").strip()
self.samples_path: Optional[Path] = Path(sp) if sp else None
self.samples: List[ExperienceSample] = []
self._embedder = None
self.embeddings: Optional[np.ndarray] = None
self.use_cache = use_cache
self.cache_path: Optional[Path] = None
self._chroma_store = None
self.use_chroma = (
use_chroma
if use_chroma is not None
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()
)
self._load_samples()
self._build_index()
if self.use_chroma:
self._init_chroma_mode()
else:
self._load_samples()
self._init_embedder_and_vector_index()
def _init_chroma_mode(self) -> None:
"""Chroma 模式:先连向量库;库非空则不再读 JSONL。"""
from utils.embedding import get_embedder
from utils.fewshot_chroma_store import FewShotChromaStore
self._embedder = get_embedder(self._embedding_model_path)
self._chroma_store = FewShotChromaStore(
self._embedder,
persist_dir=self._chroma_persist_dir,
)
n = self._chroma_store.count()
if n > 0:
logger.info(
"[Few-shot] 已从向量库加载(Chroma %s 条,%s)",
n,
self._chroma_store.persist_dir,
)
else:
self._load_samples()
if self.samples:
logger.info(
"Few-shot Chroma 库为空,正从 JSONL 写入向量索引: %s",
self.samples_path,
)
self._chroma_store.build_from_samples(self.samples, force_rebuild=False)
elif self.samples_path is None:
logger.warning(
"[Few-shot] Chroma 库为空且未配置 JSONL;请运行 "
"scripts/build_fewshot_chroma_index.py 或设置 FEWSHOT_DATA_PATH"
)
elif not self.samples_path.exists():
logger.warning(
"[Few-shot] Chroma 库为空且 JSONL 不存在: %s;请先灌库",
self.samples_path,
)
else:
logger.warning(
"[Few-shot] Chroma 库为空且 JSONL 无有效样本: %s",
self.samples_path,
)
self.embeddings = None
self.cache_path = None
logger.info("[OK] Few-shot 使用 Chroma(%s 条)", self._chroma_store.count())
def _load_samples(self):
"""加载样本数据"""
"""从 JSONL 加载样本到内存(非 Chroma 模式必需;Chroma 空库时用于首次灌库)。"""
if self.samples_path is None:
if self.use_chroma:
logger.info("[Few-shot] 未配置 JSONL 路径,运行时仅从 Chroma 检索")
return
raise ValueError("未指定样本 JSONL 路径且未启用 Chroma(FEWSHOT_USE_CHROMA)")
if not self.samples_path.exists():
if self.use_chroma:
logger.warning(
"[Few-shot] JSONL 不存在 %s,跳过文件加载,仅从 Chroma 检索",
self.samples_path,
)
return
raise FileNotFoundError(f"样本文件不存在: {self.samples_path}")
logger.info(f"加载样本: {self.samples_path}")
with self.samples_path.open('r', encoding='utf-8') as f:
logger.info("从 JSONL 加载样本: %s", self.samples_path)
with self.samples_path.open("r", encoding="utf-8") as f:
for line in f:
data = json.loads(line.strip())
self.samples.append(ExperienceSample.from_dict(data))
logger.info(f"[OK] 加载 {len(self.samples)} 个样本")
logger.info("[OK] 自 JSONL 加载 %s 个样本", len(self.samples))
def _build_index(self):
"""用项目统一 Embedder 构建语义索引"""
def _init_embedder_and_vector_index(self) -> None:
"""非 Chroma:初始化 Embedder 与内存 numpy 索引。"""
from utils.embedding import get_embedder
self._embedder = get_embedder(self._embedding_model_path)
self._build_numpy_index()
def _build_numpy_index(self) -> None:
"""内存向量 + 可选 .npy 缓存(与历史行为一致)。"""
assert self._embedder is not None
probe = self._embedder.encode(
[" "],
batch_size=1,
@@ -168,6 +249,43 @@ class FewShotSelector:
np.save(self.cache_path, self.embeddings)
logger.info(f"[OK] Few-shot 向量已缓存: {self.cache_path}")
def _select_chroma(
self,
question: str,
top_k: int,
min_rating: Optional[int],
required_tags: Optional[List[str]],
max_difficulty: str,
exclude_qids: Optional[List[str]],
) -> List[ExperienceSample]:
from utils.fewshot_chroma_store import sample_from_chroma_metadata
assert self._chroma_store is not None
over_fetch = max(top_k * 12, 48)
rows = self._chroma_store.search_raw(question, top_k=over_fetch)
candidates: List[tuple[float, ExperienceSample]] = []
for score, meta, doc in rows:
s = sample_from_chroma_metadata(meta, doc)
if exclude_qids and s.qid in exclude_qids:
continue
if min_rating and s.rating is not None and s.rating < min_rating:
continue
if max_difficulty == "easy" and s.difficulty != "easy":
continue
if max_difficulty == "medium" and s.difficulty == "hard":
continue
if required_tags and not all(tag in s.tags for tag in required_tags):
continue
candidates.append((score, s))
candidates.sort(key=lambda x: -x[0])
selected = [s for _, s in candidates[:top_k]]
logger.info(
f"Few-shot(Chroma)选择: 问题='{question[:30]}...' → 选中{len(selected)}个示例 "
f"(top_k={top_k}, min_rating={min_rating})"
)
return selected
def select(
self,
question: str,
@@ -191,6 +309,16 @@ class FewShotSelector:
Returns:
排序后的示例列表(最相关优先)
"""
if self._chroma_store is not None:
return self._select_chroma(
question,
top_k,
min_rating,
required_tags,
max_difficulty,
exclude_qids,
)
if self._embedder is None or self.embeddings is None:
logger.error("Few-shot 索引未初始化")
return []
@@ -266,10 +394,18 @@ class FewShotSelector:
return "\n".join(lines)
def _all_samples_for_aggregation(self) -> List[ExperienceSample]:
"""内存中的 JSONL 样本,或 Chroma 全量导出(用于统计/按标签列举)。"""
if self.samples:
return self.samples
if self._chroma_store is not None and self._chroma_store.count() > 0:
return self._chroma_store.get_all_as_samples()
return []
def get_tagged_examples(self, tags: List[str], top_k_per_tag: int = 2) -> str:
"""获取特定标签的示例"""
tagged_samples = []
for sample in self.samples:
for sample in self._all_samples_for_aggregation():
if any(tag in sample.tags for tag in tags):
tagged_samples.append(sample)
@@ -287,27 +423,33 @@ class FewShotSelector:
def get_stats(self) -> Dict:
"""获取数据集统计"""
src = self._all_samples_for_aggregation()
stats = {
"total": len(self.samples),
"total": len(src),
"by_rating": {},
"by_difficulty": {},
"by_tag": {},
"avg_rating": 0.0
"avg_rating": 0.0,
"source": (
"chroma"
if (not self.samples and self._chroma_store and self._chroma_store.count() > 0)
else ("jsonl" if self.samples else "empty")
),
}
ratings = [s.rating for s in self.samples if s.rating]
ratings = [s.rating for s in src if s.rating]
if ratings:
stats["avg_rating"] = sum(ratings) / len(ratings)
for r in range(1, 11):
stats["by_rating"][r] = sum(1 for s in self.samples if s.rating == r)
stats["by_rating"][r] = sum(1 for s in src if s.rating == r)
for diff in ["easy", "medium", "hard"]:
stats["by_difficulty"][diff] = sum(
1 for s in self.samples if s.difficulty == diff
1 for s in src if s.difficulty == diff
)
tag_counts = {}
for s in self.samples:
for s in src:
for tag in s.tags:
tag_counts[tag] = tag_counts.get(tag, 0) + 1
stats["by_tag"] = tag_counts