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

357 lines
12 KiB
Python
Raw Normal View History

2026-04-10 16:52:07 +08:00
"""
Few-shot示例选择器 - 基于经验数据集动态选择相关示例
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(本地 Qwen3 或远程 OpenAI 兼容 API,
由 USE_LOCAL_EMBEDDING 及 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
用法:
from utils.fewshot_selector import FewShotSelector
selector = 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__)
_DEFAULT_LOCAL_EMBED_PATH = "./data/models/Qwen3-Embedding-0.6B"
@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: str,
embedding_model_path: Optional[str] = None,
use_cache: bool = True,
):
"""
Args:
samples_path: 样本JSONL文件路径
embedding_model_path: 本地 Embedding 模型目录;None 时用环境变量
EMBEDDING_MODEL_PATH(仅 USE_LOCAL_EMBEDDING=true 时有效)
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
"""
self.samples_path = Path(samples_path)
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._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()
def _load_samples(self):
"""加载样本数据"""
if not self.samples_path.exists():
raise FileNotFoundError(f"样本文件不存在: {self.samples_path}")
logger.info(f"加载样本: {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)} 个样本")
def _build_index(self):
"""用项目统一 Embedder 构建语义索引"""
from utils.embedding import get_embedder
self._embedder = get_embedder(self._embedding_model_path)
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(
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._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 get_tagged_examples(self, tags: List[str], top_k_per_tag: int = 2) -> str:
"""获取特定标签的示例"""
tagged_samples = []
for sample in self.samples:
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:
"""获取数据集统计"""
stats = {
"total": len(self.samples),
"by_rating": {},
"by_difficulty": {},
"by_tag": {},
"avg_rating": 0.0
}
ratings = [s.rating for s in self.samples 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)
for diff in ["easy", "medium", "hard"]:
stats["by_difficulty"][diff] = sum(
1 for s in self.samples if s.difficulty == diff
)
tag_counts = {}
for s in self.samples:
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选择器"""
default_path = Path(__file__).resolve().parent.parent / "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")