Files
ai-g3sb-backman2.0/backend/utils/fewshot_selector.py
T
2026-04-14 10:28:22 +08:00

358 lines
12 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()(本地 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选择器"""
# __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")