Files
ai-g3sb-backman2.0/backend/utils/fewshot_selector.py
T

488 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Few-shot示例选择器 - 基于经验数据集动态选择相关示例
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(远程 OpenAI 兼容 API;
由 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__)
@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,
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``。
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
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._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._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")