Files
ars-opd-rebuild/scripts/generate_teacher_completions.py
T
iomgaa 404abc22bf 重构: load_sft_dataset 改吃散装参数(磨平接口回看记录的毛刺)
深模块修正:本函数只用 5 个字段,却索要整个 SFTConfig——层 1 无痛,但诊断脚本
被迫伪造 output_dir(4 处 /tmp/diag、outputs/_unused),层 2 更因 DistillConfig
无 teacher_completions_path 而无法复用。改收 dataset_path/split/subset_size/seed/
teacher_completions_path 五个散装参数(接口终于比实现轻)。

- data.py: 签名改散装参数;移除 TYPE_CHECKING 的 SFTConfig 依赖
- train_sft / diag_loss_probe / diag_collator: 仍持 SFTConfig(喂 collator),改调用点
- diag_generate / generate_teacher_completions: 只为 load 而造 config,直接丢弃、
  去掉伪造 output_dir,改传字面量
- 为 U5 层 2 训练脚本能直接 load_sft_dataset(distill_cfg 的字段) 铺路

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:21:16 -04:00

37 lines
1.8 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.
"""层 1:为 DAPO 1k 子集生成 teacherMiniMax-M3)解答缓存。
自包含实验脚本:全部参数写死在此,零参数复现。在**本地**运行(纯 API 调用,
不需要 GPU;本机可直连自建网关):
conda activate ars-opd
python -u scripts/generate_teacher_completions.py
前置:
1. .env 已填 TEACHER_API_BASE / TEACHER_API_KEY / TEACHER_MODEL
2. DAPO parquet 已下载到 DATASET_PATH(见 docs/02 §4)。
中断安全:缓存逐条落盘,重跑本脚本自动跳过已完成条目(断点续传)。
"""
from ars_opd.configs import TeacherGenConfig
from ars_opd.data import load_sft_dataset
from ars_opd.teacher import TeacherClient, generate_completions
# 非显然约束:这里的 dataset/subset_size/seed 必须与 T5 训练脚本完全一致——
# 两侧各自走"加载→归一→抽子集",seed 相同才是同一批题(data.py 有详注)
DATASET_PATH = "data/dapo-math-17k-unique.parquet" # DAPO 官方去重版,17917 行
CACHE_PATH = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
# 试跑说明:首次建议把下面 subset_size 临时改成 5,跑通并人工抽查缓存里的解答
# 质量(think 是否剥净、格式是否正常)后再改回 1000 重跑。放心改:subset 是对
# 同一 seed 的洗牌序列取前缀,前 5 条与前 1000 条的头 5 条完全相同,试跑写入的
# 缓存在正式跑时全部命中,一分钱不浪费。
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42)
prompts = [row["messages"] for row in dataset]
teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认
generate_completions(prompts, CACHE_PATH, teacher)
print(f"完成。缓存文件:{CACHE_PATH}")