"""实验配置(当前只含层 1 的 SFTConfig,后续层在此文件追加各自的 dataclass)。 设计约定(对应 CLAUDE.md §2"配置显式化"): - 所有实验参数必须是这里某个 dataclass 的字段;代码里出现魔法数字/路径即违规。 - 密钥(API key 等)不进配置类,走 `.env`(见 teacher.py)。 - 机器相关路径(数据集、输出目录)不给默认值,强制调用方显式传入—— 防止参考实现里 `/fsx` 硬编码那类"在别人机器上必炸"的坑。 """ from __future__ import annotations from dataclasses import dataclass @dataclass(frozen=True) class SFTConfig: """层 1 SFT 基线的全部实验参数。 论文锚点:§3.1 式(1) 的标准交叉熵 SFT;但按 §5.1 的基线定义, 训练数据是 teacher rollout(离线蒸馏),不是人写答案——所以有 `teacher_completions_path` 字段:DAPO 是 prompt-only 数据集, 解答一律来自 teacher 生成的缓存文件。 frozen=True:配置一旦构造即只读。训练中途被悄悄改掉的配置是最难 排查的 bug 来源之一;要换参数就构造一个新实例,留下明确的代码痕迹。 """ # ---- 机器相关路径(无默认值,必须显式传入)---- dataset_path: str """DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。""" output_dir: str """checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。""" # ---- 数据 ---- teacher_completions_path: str | None = None """teacher 解答缓存(teacher.py 生成的 JSONL)。None 表示数据集自带 assistant 轮次;若数据实际是 prompt-only 又没给此路径,data.py 会显式报错, 不做静默兜底。""" dataset_split: str = "train" subset_size: int | None = 1000 """随机抽取的子集大小(控制 teacher API 成本,roadmap 定为 ~1k);None = 全量。""" # ---- 序列双预算 ---- # 非显然约束:prompt 与 completion 必须各有独立预算。若只用一个 max_length # 从右截断,超长解答会把 prompt 挤空,模型在"没有题目"的样本上学习解答 # ——这是参考实现 collator(trainer:267-292) 的头号正确性卖点,此处继承。 max_length: int = 4096 """prompt + completion 的总 token 预算。""" max_prompt_length: int = 1024 """prompt 单独预算;completion 实际预算 = max_length - len(截断后 prompt)。""" enable_thinking: bool = False """Qwen3 chat 模板的思考开关。False 时模板注入空 `\\n\\n`。 非显然约束:此开关改变渲染后的 prompt 文本,从而改变 prompt/completion 的 token 边界——训练与推理必须取同一值,否则掩码整体错位。""" # ---- 优化 ---- # 差异标注:论文 §5.1 的蒸馏训练用 lr=1e-6,参考实现 SFT 默认 2e-5 # (train_distillation.py:73)。SFT 有真实 token 监督、信号密集,从参考实现取 # 2e-5;层 5 的蒸馏配置再回到论文的 1e-6。 learning_rate: float = 2e-5 per_device_train_batch_size: int = 2 """非显然约束:别看 0.6B 小就调大它——显存大头是 (B,T,V) 的 logits 链 (fp32 一份 ~20G@B=8)与逐层激活,都正比于 B 而与参数量无关;B=8 实测 爆 80G 卡(2026-07-18 远程 sanity)。""" gradient_accumulation_steps: int = 8 """全局 batch = 2(per_device) × 4(卡) × 8(累积) = 64,与参考实现注释的 训练规模(trainer 配置注释"global batch 64")对齐。""" num_train_epochs: int = 1 max_steps: int = -1 """>0 时覆盖 num_train_epochs,只跑这么多步——远程 50 步 sanity 用;-1 = 按 epoch。""" lr_scheduler_type: str = "linear" warmup_ratio: float = 0.0 gradient_checkpointing: bool = False """0.6B 学生显存富余,不开(省 ~40% 显存、慢 ~30%)。注意:FSDP 下此开关 是 no-op,真正的开关是 FSDP_ACTIVATION_CHECKPOINTING 环境变量(见 scripts/ 训练脚本头部的前置块,docs/02 §2.6)。""" bf16: bool = True seed: int = 42 # ---- 日志与保存 ---- logging_steps: int = 1 save_steps: int = 100 save_total_limit: int = 2 report_to: str = "none" """"none" 或 "wandb"。默认 none:本地调试不该悄悄往外发数据,远程脚本显式开。""" def __post_init__(self) -> None: """构造即校验:配置错误必须在训练开始前炸,而不是跑到第一个超长样本才炸。""" if self.max_prompt_length >= self.max_length: raise ValueError( f"max_prompt_length({self.max_prompt_length}) 必须小于 " f"max_length({self.max_length}),否则 completion 预算为零," f"所有样本的 labels 将全为 -100,loss 恒为 0 且无报错——静默空训练。" ) if self.learning_rate <= 0: raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}") if self.subset_size is not None and self.subset_size <= 0: raise ValueError( f"subset_size 必须为正整数或 None(全量),收到 {self.subset_size}" ) if self.max_steps == 0 or self.max_steps < -1: raise ValueError( f"max_steps 只接受 -1(按 epoch)或正整数,收到 {self.max_steps}" ) @dataclass(frozen=True) class TeacherGenConfig: """teacher 批量生成(层 1 能力)的采样与执行参数。 连接信息(API 地址/密钥/模型名)不在这里——那是部署环境的事实,走 `.env` (teacher.py 读取);这里只放"换一组值就是换一个实验"的采样参数。 """ temperature: float = 1.0 top_p: float = 0.95 """MiniMax M 系官方推荐采样参数:temperature=1.0, top_p=0.95。""" max_tokens: int = 16384 """teacher 单条回复的 token 上限。这是上限不是目标——按实际生成量计费, 放大它不增加正常解答的成本,只给最难的题留出写完的空间(8192 时 59 条实测 截断 2 条)。非显然约束:M3 的思考段也计入此额度,设太小会把解答挤没。""" strip_think: bool = True """剥离 content 开头的 ... 思考段。SFT 的监督目标是最终 解答;student 以 enable_thinking=False 训练,学思考段会与模板约定矛盾。""" concurrency: int = 16 """并发请求数(线程池大小)。上限看网关的承受力,报 429 就调小。""" max_retries: int = 3 """单请求的网络级重试次数(openai 客户端内建指数退避)。""" system_prompt: str | None = None """None = 不加 system 轮(DAPO 题面自带作答指令,不需要额外指挥)。""" def __post_init__(self) -> None: if self.max_tokens <= 0: raise ValueError(f"max_tokens 必须为正,收到 {self.max_tokens}") if self.concurrency < 1: raise ValueError(f"concurrency 必须 ≥1,收到 {self.concurrency}") if self.temperature < 0: raise ValueError(f"temperature 必须 ≥0,收到 {self.temperature}")