2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
Few-shot示例选择器 - 基于经验数据集动态选择相关示例
|
|
|
|
|
|
|
2026-04-15 09:49:18 +08:00
|
|
|
|
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(远程 OpenAI 兼容 API;
|
|
|
|
|
|
由 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
用法:
|
|
|
|
|
|
from utils.fewshot_selector import FewShotSelector
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
# Chroma 模式(FEWSHOT_USE_CHROMA=true):可不传 JSONL,路径传 None 或 ""
|
|
|
|
|
|
selector = FewShotSelector(None)
|
|
|
|
|
|
# 或传统:FewShotSelector("data/experiences/all_samples.jsonl")
|
2026-04-10 16:52:07 +08:00
|
|
|
|
examples = selector.select(question="查询2024年1月的销售额", top_k=3)
|
|
|
|
|
|
|
|
|
|
|
|
# 在Prompt中使用
|
|
|
|
|
|
prompt = f"{examples}\n当前问题:{question}\nSchema:{schema}"
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
import json
|
|
|
|
|
|
import os
|
|
|
|
|
|
from pathlib import Path
|
2026-04-16 13:48:44 +08:00
|
|
|
|
from typing import List, Dict, Optional, Tuple
|
2026-04-10 16:52:07 +08:00
|
|
|
|
from dataclasses import dataclass
|
|
|
|
|
|
import numpy as np
|
|
|
|
|
|
import logging
|
|
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
|
class ExperienceSample:
|
|
|
|
|
|
"""经验数据样本"""
|
|
|
|
|
|
qid: str
|
|
|
|
|
|
question_zh: str
|
|
|
|
|
|
question_en: Optional[str]
|
|
|
|
|
|
sql: str
|
|
|
|
|
|
explanation: str
|
|
|
|
|
|
rating: Optional[int]
|
|
|
|
|
|
tags: List[str]
|
|
|
|
|
|
difficulty: str
|
|
|
|
|
|
|
|
|
|
|
|
@classmethod
|
|
|
|
|
|
def from_dict(cls, data: dict) -> "ExperienceSample":
|
|
|
|
|
|
return cls(
|
|
|
|
|
|
qid=data.get("qid", ""),
|
|
|
|
|
|
question_zh=data.get("question_zh", ""),
|
|
|
|
|
|
question_en=data.get("question_en"),
|
|
|
|
|
|
sql=data.get("sql", ""),
|
|
|
|
|
|
explanation=data.get("explanation", ""),
|
|
|
|
|
|
rating=data.get("rating"),
|
|
|
|
|
|
tags=data.get("tags", []),
|
|
|
|
|
|
difficulty=data.get("difficulty", "medium")
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
def to_fewshot_format(self, include_explanation: bool = True) -> str:
|
|
|
|
|
|
"""转换为few-shot格式"""
|
|
|
|
|
|
result = f"问题:{self.question_zh}\nSQL:\n{self.sql}"
|
|
|
|
|
|
if include_explanation and self.explanation:
|
|
|
|
|
|
result += f"\n说明:{self.explanation[:200]}"
|
|
|
|
|
|
return result
|
|
|
|
|
|
|
|
|
|
|
|
def to_dict(self) -> dict:
|
|
|
|
|
|
return {
|
|
|
|
|
|
"qid": self.qid,
|
|
|
|
|
|
"question": self.question_zh,
|
|
|
|
|
|
"sql": self.sql,
|
|
|
|
|
|
"rating": self.rating,
|
|
|
|
|
|
"tags": self.tags,
|
|
|
|
|
|
"difficulty": self.difficulty
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class FewShotSelector:
|
|
|
|
|
|
"""
|
|
|
|
|
|
Few-shot示例选择器
|
|
|
|
|
|
|
|
|
|
|
|
根据用户问题,从经验数据集中检索最相似的示例,
|
|
|
|
|
|
用于增强Prompt,提升LLM生成质量。
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
def __init__(
|
|
|
|
|
|
self,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
samples_path: Optional[str] = None,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
use_cache: bool = True,
|
2026-04-14 18:02:12 +08:00
|
|
|
|
use_chroma: Optional[bool] = None,
|
|
|
|
|
|
chroma_persist_dir: Optional[str] = None,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
):
|
|
|
|
|
|
"""
|
|
|
|
|
|
Args:
|
2026-04-14 18:02:12 +08:00
|
|
|
|
samples_path: 样本 JSONL 路径;**空字符串**表示不读文件(仅当 ``FEWSHOT_USE_CHROMA=true``
|
|
|
|
|
|
且 Chroma 中已有数据时可用)。构建索引仍请用 ``scripts/build_fewshot_chroma_index.py --samples``。
|
2026-04-10 16:52:07 +08:00
|
|
|
|
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
|
2026-04-14 18:02:12 +08:00
|
|
|
|
use_chroma: 是否使用 Chroma 持久化向量库;None 时读环境变量 FEWSHOT_USE_CHROMA
|
|
|
|
|
|
chroma_persist_dir: Chroma 目录;None 时用 FEWSHOT_CHROMA_PATH 或默认 chroma_fewshot
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
2026-04-14 18:02:12 +08:00
|
|
|
|
sp = (samples_path or "").strip()
|
|
|
|
|
|
self.samples_path: Optional[Path] = Path(sp) if sp else None
|
2026-04-10 16:52:07 +08:00
|
|
|
|
self.samples: List[ExperienceSample] = []
|
|
|
|
|
|
self._embedder = None
|
|
|
|
|
|
self.embeddings: Optional[np.ndarray] = None
|
|
|
|
|
|
self.use_cache = use_cache
|
|
|
|
|
|
self.cache_path: Optional[Path] = None
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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
|
|
|
|
|
|
|
2026-04-15 09:49:18 +08:00
|
|
|
|
self._embedder = get_embedder()
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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())
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-16 13:48:44 +08:00
|
|
|
|
@property
|
|
|
|
|
|
def is_chroma_backend(self) -> bool:
|
|
|
|
|
|
"""是否使用 Chroma 持久化库(`data/embeddings/chroma_fewshot` 等)。"""
|
|
|
|
|
|
return self._chroma_store is not None
|
|
|
|
|
|
|
|
|
|
|
|
def _passes_filters(
|
|
|
|
|
|
self,
|
|
|
|
|
|
sample: "ExperienceSample",
|
|
|
|
|
|
min_rating: Optional[int],
|
|
|
|
|
|
required_tags: Optional[List[str]],
|
|
|
|
|
|
max_difficulty: str,
|
|
|
|
|
|
exclude_qids: Optional[List[str]],
|
|
|
|
|
|
) -> bool:
|
|
|
|
|
|
if exclude_qids and sample.qid in exclude_qids:
|
|
|
|
|
|
return False
|
|
|
|
|
|
if min_rating and sample.rating is not None and sample.rating < min_rating:
|
|
|
|
|
|
return False
|
|
|
|
|
|
if max_difficulty == "easy" and sample.difficulty != "easy":
|
|
|
|
|
|
return False
|
|
|
|
|
|
if max_difficulty == "medium" and sample.difficulty == "hard":
|
|
|
|
|
|
return False
|
|
|
|
|
|
if required_tags and not all(tag in sample.tags for tag in required_tags):
|
|
|
|
|
|
return False
|
|
|
|
|
|
return True
|
|
|
|
|
|
|
|
|
|
|
|
def select_best_with_score(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
min_rating: Optional[int] = None,
|
|
|
|
|
|
required_tags: Optional[List[str]] = None,
|
|
|
|
|
|
max_difficulty: str = "hard",
|
|
|
|
|
|
exclude_qids: Optional[List[str]] = None,
|
|
|
|
|
|
) -> Optional[Tuple[ExperienceSample, float]]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
返回通过筛选的**相似度最高**一条样本及分数 ``[0,1]``(与向量余弦一致:1 - distance)。
|
|
|
|
|
|
无命中时返回 ``None``。
|
|
|
|
|
|
"""
|
|
|
|
|
|
if self._chroma_store is not None:
|
|
|
|
|
|
from utils.fewshot_chroma_store import sample_from_chroma_metadata
|
|
|
|
|
|
|
|
|
|
|
|
over_fetch = 96
|
|
|
|
|
|
rows = self._chroma_store.search_raw(question, top_k=over_fetch)
|
|
|
|
|
|
for score, meta, doc in rows:
|
|
|
|
|
|
s = sample_from_chroma_metadata(meta, doc)
|
|
|
|
|
|
if not self._passes_filters(
|
|
|
|
|
|
s, min_rating, required_tags, max_difficulty, exclude_qids
|
|
|
|
|
|
):
|
|
|
|
|
|
continue
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"[Few-shot] best_with_score: qid=%s score=%.4f preview=%r",
|
|
|
|
|
|
s.qid,
|
|
|
|
|
|
float(score),
|
|
|
|
|
|
(s.question_zh or "")[:100],
|
|
|
|
|
|
)
|
|
|
|
|
|
return (s, float(score))
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
if self._embedder is None or self.embeddings is None or len(self.samples) == 0:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
q_emb = self._embedder.encode(
|
|
|
|
|
|
[question],
|
|
|
|
|
|
batch_size=1,
|
|
|
|
|
|
normalize=True,
|
|
|
|
|
|
show_progress=False,
|
|
|
|
|
|
)[0]
|
|
|
|
|
|
scores = np.dot(self.embeddings, q_emb)
|
|
|
|
|
|
best: Optional[Tuple[float, ExperienceSample]] = None
|
|
|
|
|
|
for idx, (score, sample) in enumerate(zip(scores, self.samples)):
|
|
|
|
|
|
if not self._passes_filters(
|
|
|
|
|
|
sample, min_rating, required_tags, max_difficulty, exclude_qids
|
|
|
|
|
|
):
|
|
|
|
|
|
continue
|
|
|
|
|
|
s = float(score)
|
|
|
|
|
|
if best is None or s > best[0]:
|
|
|
|
|
|
best = (s, sample)
|
|
|
|
|
|
if best is None:
|
|
|
|
|
|
return None
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"[Few-shot] best_with_score(numpy): qid=%s score=%.4f preview=%r",
|
|
|
|
|
|
best[1].qid,
|
|
|
|
|
|
best[0],
|
|
|
|
|
|
(best[1].question_zh or "")[:100],
|
|
|
|
|
|
)
|
|
|
|
|
|
return (best[1], best[0])
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def _load_samples(self):
|
2026-04-14 18:02:12 +08:00
|
|
|
|
"""从 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)")
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if not self.samples_path.exists():
|
2026-04-14 18:02:12 +08:00
|
|
|
|
if self.use_chroma:
|
|
|
|
|
|
logger.warning(
|
|
|
|
|
|
"[Few-shot] JSONL 不存在 %s,跳过文件加载,仅从 Chroma 检索",
|
|
|
|
|
|
self.samples_path,
|
|
|
|
|
|
)
|
|
|
|
|
|
return
|
2026-04-10 16:52:07 +08:00
|
|
|
|
raise FileNotFoundError(f"样本文件不存在: {self.samples_path}")
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
logger.info("从 JSONL 加载样本: %s", self.samples_path)
|
|
|
|
|
|
with self.samples_path.open("r", encoding="utf-8") as f:
|
2026-04-10 16:52:07 +08:00
|
|
|
|
for line in f:
|
|
|
|
|
|
data = json.loads(line.strip())
|
|
|
|
|
|
self.samples.append(ExperienceSample.from_dict(data))
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
logger.info("[OK] 自 JSONL 加载 %s 个样本", len(self.samples))
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
def _init_embedder_and_vector_index(self) -> None:
|
|
|
|
|
|
"""非 Chroma:初始化 Embedder 与内存 numpy 索引。"""
|
2026-04-10 16:52:07 +08:00
|
|
|
|
from utils.embedding import get_embedder
|
|
|
|
|
|
|
2026-04-15 09:49:18 +08:00
|
|
|
|
self._embedder = get_embedder()
|
2026-04-14 18:02:12 +08:00
|
|
|
|
self._build_numpy_index()
|
|
|
|
|
|
|
|
|
|
|
|
def _build_numpy_index(self) -> None:
|
|
|
|
|
|
"""内存向量 + 可选 .npy 缓存(与历史行为一致)。"""
|
|
|
|
|
|
assert self._embedder is not None
|
2026-04-10 16:52:07 +08:00
|
|
|
|
probe = self._embedder.encode(
|
|
|
|
|
|
[" "],
|
|
|
|
|
|
batch_size=1,
|
|
|
|
|
|
normalize=True,
|
|
|
|
|
|
show_progress=False,
|
|
|
|
|
|
)
|
|
|
|
|
|
dim = int(probe.shape[1])
|
|
|
|
|
|
self.cache_path = (
|
|
|
|
|
|
self.samples_path.parent / f"{self.samples_path.stem}.fewshot_dim{dim}.npy"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if self.use_cache and self.cache_path.exists():
|
|
|
|
|
|
try:
|
|
|
|
|
|
self.embeddings = np.load(self.cache_path)
|
|
|
|
|
|
if (
|
|
|
|
|
|
self.embeddings.shape[0] == len(self.samples)
|
|
|
|
|
|
and self.embeddings.shape[1] == dim
|
|
|
|
|
|
):
|
|
|
|
|
|
logger.info(f"[OK] 加载 Few-shot 向量缓存: {self.cache_path}")
|
|
|
|
|
|
return
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.warning(f"Few-shot 缓存加载失败: {e},将重新计算")
|
|
|
|
|
|
|
|
|
|
|
|
# 远程 API 通常不接受空字符串作 input
|
|
|
|
|
|
questions = [
|
|
|
|
|
|
(s.question_zh or "").strip() or " "
|
|
|
|
|
|
for s in self.samples
|
|
|
|
|
|
]
|
|
|
|
|
|
if not questions:
|
|
|
|
|
|
self.embeddings = np.empty((0, dim), dtype=np.float32)
|
|
|
|
|
|
return
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(f"计算 {len(questions)} 个 Few-shot 样本向量...")
|
|
|
|
|
|
self.embeddings = self._embedder.encode(
|
|
|
|
|
|
questions,
|
|
|
|
|
|
batch_size=min(32, len(questions)),
|
|
|
|
|
|
normalize=True,
|
|
|
|
|
|
show_progress=True,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
if self.use_cache and self.cache_path is not None:
|
|
|
|
|
|
np.save(self.cache_path, self.embeddings)
|
|
|
|
|
|
logger.info(f"[OK] Few-shot 向量已缓存: {self.cache_path}")
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def select(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
top_k: int = 3,
|
|
|
|
|
|
min_rating: Optional[int] = None,
|
|
|
|
|
|
required_tags: Optional[List[str]] = None,
|
|
|
|
|
|
max_difficulty: str = "hard",
|
|
|
|
|
|
exclude_qids: Optional[List[str]] = None
|
|
|
|
|
|
) -> List[ExperienceSample]:
|
|
|
|
|
|
"""
|
|
|
|
|
|
选择最相关的few-shot示例
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
question: 用户问题
|
|
|
|
|
|
top_k: 返回示例数量
|
|
|
|
|
|
min_rating: 最低评分(None表示不限制)
|
|
|
|
|
|
required_tags: 必须包含的标签(如["aggregation", "join"])
|
|
|
|
|
|
max_difficulty: 最大难度(过滤更难的示例)
|
|
|
|
|
|
exclude_qids: 排除的QID(避免与当前问题相同)
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
排序后的示例列表(最相关优先)
|
|
|
|
|
|
"""
|
2026-04-14 18:02:12 +08:00
|
|
|
|
if self._chroma_store is not None:
|
|
|
|
|
|
return self._select_chroma(
|
|
|
|
|
|
question,
|
|
|
|
|
|
top_k,
|
|
|
|
|
|
min_rating,
|
|
|
|
|
|
required_tags,
|
|
|
|
|
|
max_difficulty,
|
|
|
|
|
|
exclude_qids,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if self._embedder is None or self.embeddings is None:
|
|
|
|
|
|
logger.error("Few-shot 索引未初始化")
|
|
|
|
|
|
return []
|
|
|
|
|
|
|
|
|
|
|
|
if len(self.samples) == 0:
|
|
|
|
|
|
return []
|
|
|
|
|
|
|
|
|
|
|
|
q_emb = self._embedder.encode(
|
|
|
|
|
|
[question],
|
|
|
|
|
|
batch_size=1,
|
|
|
|
|
|
normalize=True,
|
|
|
|
|
|
show_progress=False,
|
|
|
|
|
|
)[0]
|
|
|
|
|
|
|
|
|
|
|
|
scores = np.dot(self.embeddings, q_emb)
|
|
|
|
|
|
|
|
|
|
|
|
candidates = []
|
|
|
|
|
|
for idx, (score, sample) in enumerate(zip(scores, self.samples)):
|
|
|
|
|
|
if exclude_qids and sample.qid in exclude_qids:
|
|
|
|
|
|
continue
|
|
|
|
|
|
if min_rating and sample.rating and sample.rating < min_rating:
|
|
|
|
|
|
continue
|
|
|
|
|
|
if max_difficulty == "easy" and sample.difficulty != "easy":
|
|
|
|
|
|
continue
|
|
|
|
|
|
if max_difficulty == "medium" and sample.difficulty == "hard":
|
|
|
|
|
|
continue
|
|
|
|
|
|
if required_tags and not all(tag in sample.tags for tag in required_tags):
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
candidates.append((idx, score, sample))
|
|
|
|
|
|
|
|
|
|
|
|
candidates.sort(key=lambda x: -x[1])
|
|
|
|
|
|
|
|
|
|
|
|
selected = [sample for _, _, sample in candidates[:top_k]]
|
|
|
|
|
|
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
f"Few-shot选择: 问题='{question[:30]}...' "
|
|
|
|
|
|
f"→ 选中{len(selected)}个示例 (top_k={top_k}, min_rating={min_rating})"
|
|
|
|
|
|
)
|
|
|
|
|
|
for s in selected:
|
|
|
|
|
|
logger.debug(
|
|
|
|
|
|
f" [{s.qid}] {s.question_zh[:50]}... (rating={s.rating}, tags={s.tags[:3]})"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
return selected
|
|
|
|
|
|
|
|
|
|
|
|
def get_examples_prompt(
|
|
|
|
|
|
self,
|
|
|
|
|
|
question: str,
|
|
|
|
|
|
top_k: int = 3,
|
|
|
|
|
|
min_rating: int = 7,
|
|
|
|
|
|
**kwargs
|
|
|
|
|
|
) -> str:
|
|
|
|
|
|
"""
|
|
|
|
|
|
生成few-shot prompt片段
|
|
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
格式化的示例字符串,可直接插入Prompt
|
|
|
|
|
|
"""
|
|
|
|
|
|
examples = self.select(question, top_k=top_k, min_rating=min_rating, **kwargs)
|
|
|
|
|
|
|
|
|
|
|
|
if not examples:
|
|
|
|
|
|
return ""
|
|
|
|
|
|
|
|
|
|
|
|
lines = ["以下为相似问题的参考SQL示例:\n"]
|
|
|
|
|
|
for i, ex in enumerate(examples, 1):
|
|
|
|
|
|
lines.append(f"示例{i}:")
|
|
|
|
|
|
lines.append(f"问题:{ex.question_zh}")
|
|
|
|
|
|
lines.append(f"SQL:\n{ex.sql}")
|
|
|
|
|
|
if ex.explanation:
|
|
|
|
|
|
lines.append(f"说明:{ex.explanation[:150]}...")
|
|
|
|
|
|
lines.append("") # 空行分隔
|
|
|
|
|
|
|
|
|
|
|
|
return "\n".join(lines)
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
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 []
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def get_tagged_examples(self, tags: List[str], top_k_per_tag: int = 2) -> str:
|
|
|
|
|
|
"""获取特定标签的示例"""
|
|
|
|
|
|
tagged_samples = []
|
2026-04-14 18:02:12 +08:00
|
|
|
|
for sample in self._all_samples_for_aggregation():
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if any(tag in sample.tags for tag in tags):
|
|
|
|
|
|
tagged_samples.append(sample)
|
|
|
|
|
|
|
|
|
|
|
|
tagged_samples.sort(key=lambda s: -(s.rating or 0))
|
|
|
|
|
|
selected = tagged_samples[:top_k_per_tag * len(tags)]
|
|
|
|
|
|
|
|
|
|
|
|
lines = [f"# {tags} 相关示例\n"]
|
|
|
|
|
|
for ex in selected:
|
|
|
|
|
|
lines.append(f"## {ex.qid}. {ex.question_zh[:50]}")
|
|
|
|
|
|
lines.append(f"评分: {ex.rating}/10")
|
|
|
|
|
|
lines.append(f"标签: {', '.join(ex.tags)}")
|
2026-04-16 16:40:22 +08:00
|
|
|
|
lines.append("SQL(纯文本,允许 -- 注释;禁止 ``` 围栏):")
|
|
|
|
|
|
lines.append(ex.sql)
|
|
|
|
|
|
lines.append("")
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
return "\n".join(lines)
|
|
|
|
|
|
|
|
|
|
|
|
def get_stats(self) -> Dict:
|
|
|
|
|
|
"""获取数据集统计"""
|
2026-04-14 18:02:12 +08:00
|
|
|
|
src = self._all_samples_for_aggregation()
|
2026-04-10 16:52:07 +08:00
|
|
|
|
stats = {
|
2026-04-14 18:02:12 +08:00
|
|
|
|
"total": len(src),
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"by_rating": {},
|
|
|
|
|
|
"by_difficulty": {},
|
|
|
|
|
|
"by_tag": {},
|
2026-04-14 18:02:12 +08:00
|
|
|
|
"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")
|
|
|
|
|
|
),
|
2026-04-10 16:52:07 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
ratings = [s.rating for s in src if s.rating]
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if ratings:
|
|
|
|
|
|
stats["avg_rating"] = sum(ratings) / len(ratings)
|
|
|
|
|
|
for r in range(1, 11):
|
2026-04-14 18:02:12 +08:00
|
|
|
|
stats["by_rating"][r] = sum(1 for s in src if s.rating == r)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
for diff in ["easy", "medium", "hard"]:
|
|
|
|
|
|
stats["by_difficulty"][diff] = sum(
|
2026-04-14 18:02:12 +08:00
|
|
|
|
1 for s in src if s.difficulty == diff
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
tag_counts = {}
|
2026-04-14 18:02:12 +08:00
|
|
|
|
for s in src:
|
2026-04-10 16:52:07 +08:00
|
|
|
|
for tag in s.tags:
|
|
|
|
|
|
tag_counts[tag] = tag_counts.get(tag, 0) + 1
|
|
|
|
|
|
stats["by_tag"] = tag_counts
|
|
|
|
|
|
|
|
|
|
|
|
return stats
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def load_fewshot_selector() -> FewShotSelector:
|
|
|
|
|
|
"""加载默认的few-shot选择器"""
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# __file__ = backend/utils/fewshot_selector.py → 仓库根为 parents[2]
|
|
|
|
|
|
default_path = Path(__file__).resolve().parents[2] / "data" / "experiences" / "all_samples.jsonl"
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return FewShotSelector(str(default_path))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
|
import argparse
|
|
|
|
|
|
|
|
|
|
|
|
parser = argparse.ArgumentParser(description="Few-shot示例选择器")
|
|
|
|
|
|
parser.add_argument("--samples", default="data/experiences/all_samples.jsonl")
|
|
|
|
|
|
parser.add_argument("--question", help="测试问题")
|
|
|
|
|
|
parser.add_argument("--top-k", type=int, default=3)
|
|
|
|
|
|
parser.add_argument("--min-rating", type=int, default=7)
|
|
|
|
|
|
parser.add_argument("--stats", action="store_true", help="显示数据集统计")
|
|
|
|
|
|
|
|
|
|
|
|
args = parser.parse_args()
|
|
|
|
|
|
|
|
|
|
|
|
selector = FewShotSelector(args.samples)
|
|
|
|
|
|
|
|
|
|
|
|
if args.stats:
|
|
|
|
|
|
stats = selector.get_stats()
|
|
|
|
|
|
print("📊 数据集统计:")
|
|
|
|
|
|
print(f" 总样本: {stats['total']}")
|
|
|
|
|
|
print(f" 平均评分: {stats['avg_rating']:.1f}")
|
|
|
|
|
|
print(f" 难度分布: {stats['by_difficulty']}")
|
|
|
|
|
|
print(f"\n Top 10 标签:")
|
|
|
|
|
|
sorted_tags = sorted(stats["by_tag"].items(), key=lambda x: -x[1])[:10]
|
|
|
|
|
|
for tag, count in sorted_tags:
|
|
|
|
|
|
print(f" {tag}: {count}")
|
|
|
|
|
|
elif args.question:
|
|
|
|
|
|
examples = selector.select(args.question, top_k=args.top_k, min_rating=args.min_rating)
|
|
|
|
|
|
print(f"\n为问题 '{args.question}' 选择的示例:\n")
|
|
|
|
|
|
for ex in examples:
|
|
|
|
|
|
print(f"[{ex.qid}] 评分:{ex.rating} 难度:{ex.difficulty}")
|
|
|
|
|
|
print(f"问题: {ex.question_zh}")
|
|
|
|
|
|
print(f"SQL:\n{ex.sql}\n")
|
|
|
|
|
|
else:
|
|
|
|
|
|
print("请指定 --question 或 --stats")
|