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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user