""" Few-shot示例选择器 - 基于经验数据集动态选择相关示例 与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(本地 Qwen3 或远程 OpenAI 兼容 API, 由 USE_LOCAL_EMBEDDING 及 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。 用法: from utils.fewshot_selector import FewShotSelector # 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中使用 prompt = f"{examples}\n当前问题:{question}\nSchema:{schema}" """ import json import os from pathlib import Path from typing import List, Dict, Optional from dataclasses import dataclass import numpy as np import logging logger = logging.getLogger(__name__) # 仅 USE_LOCAL_EMBEDDING=true 时通过 EMBEDDING_MODEL_PATH 使用;远程模式留空即可 _DEFAULT_LOCAL_EMBED_PATH = "" @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, 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 路径;**空字符串**表示不读文件(仅当 ``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 """ 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() ) 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("从 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("[OK] 自 JSONL 加载 %s 个样本", len(self.samples)) 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, 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}") 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, 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: 排序后的示例列表(最相关优先) """ 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 [] 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) 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._all_samples_for_aggregation(): 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)}") lines.append(f"```sql\n{ex.sql}\n```\n") return "\n".join(lines) def get_stats(self) -> Dict: """获取数据集统计""" src = self._all_samples_for_aggregation() stats = { "total": len(src), "by_rating": {}, "by_difficulty": {}, "by_tag": {}, "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 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 src if s.rating == r) for diff in ["easy", "medium", "hard"]: stats["by_difficulty"][diff] = sum( 1 for s in src if s.difficulty == diff ) tag_counts = {} for s in src: 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选择器""" # __file__ = backend/utils/fewshot_selector.py → 仓库根为 parents[2] default_path = Path(__file__).resolve().parents[2] / "data" / "experiences" / "all_samples.jsonl" 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")