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
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