8b362eae09
首冒烟发现:sanity 的 loss 平滑、无预期毛刺,因 HF 默认 max_grad_norm=1.0 把 反向 KL 的梯度爆炸(§4.1,实测 grad_norm 14→2 是裁剪前范数)默默压平了——正是 本项目要堵的"静默行为"。 - configs.py: DistillConfig 加 max_grad_norm=1.0(默认=原 HF 行为),docstring 讲清 它是 §4.1 爆炸的隐形稳定器、日志 grad_norm 是裁剪前值;__post_init__ 校验 >0 - train_whitebox.py: FULL 显式写出、TrainingArguments 传入;build_config 加 noclip 模式(max_grad_norm=1e9≈关裁剪 + lr 5× + 15 步)暴露原始爆炸供教学对照 - .sh: 用法加 noclip 模式说明 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
311 lines
16 KiB
Python
311 lines
16 KiB
Python
"""实验配置(层 1:SFTConfig / TeacherGenConfig;层 2:DistillConfig)。
|
||
|
||
设计约定(对应 CLAUDE.md §2"配置显式化"):
|
||
- 所有实验参数必须是这里某个 dataclass 的字段;代码里出现魔法数字/路径即违规。
|
||
- 密钥(API key 等)不进配置类,走 `.env`(见 teacher.py)。
|
||
- 机器相关路径(数据集、输出目录)不给默认值,强制调用方显式传入——
|
||
防止参考实现里 `/fsx` 硬编码那类"在别人机器上必炸"的坑。
|
||
- 各层的 config 自包含、不互相继承:层与层是不同实验,共享基类会把它们耦合,
|
||
违背"从上读到下看懂全部流程"(CLAUDE.md §2)。字段重复是有意接受的成本。
|
||
"""
|
||
|
||
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 时模板注入空 `<think>\\n\\n</think>`。
|
||
非显然约束:此开关改变渲染后的 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 开头的 <think>...</think> 思考段。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}")
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class DistillConfig:
|
||
"""层 2 white-box OPD 基线(token 级反向 KL 蒸馏)的全部实验参数。
|
||
|
||
论文锚点:§3.1 式(2) 的 on-policy 白盒蒸馏 L = E_{y~π_θ}[Σ_t KL(π_θ ‖ π_T)]。
|
||
与层 1 SFTConfig 的三处结构性差异(docs/03 §3 偏差清单):
|
||
- 无 teacher_completions_path:teacher 现场前向给出全词表 logits、student 现场
|
||
on-policy 生成轨迹,两者都不落盘缓存,故层 2 不需要 teacher 解答文件。
|
||
- 无 max_length(总预算):completion 不再来自数据,而是 model.generate 生成,
|
||
序列总长 = prompt(≤max_prompt_length) + 生成(≤max_new_tokens),由两个预算
|
||
各自界定,不需要一个总的右截断预算。
|
||
- 无 top_k:§4 删除清单——本地同 tokenizer teacher 放得下全词表,恒走精确
|
||
全词表 KL,不做参考实现默认的 top-1 稀疏近似(那是 API 传输妥协,非论文成分)。
|
||
|
||
teacher_model 在此、student 在脚本(U5 的常量,同层 1 的 STUDENT_MODEL):
|
||
student 是被训练的固定基线,teacher 是"换一个就是换一个实验"的旋钮,故归 config。
|
||
"""
|
||
|
||
# ---- 机器相关路径(无默认值,必须显式传入)----
|
||
dataset_path: str
|
||
"""DAPO-Math-17K 的本地 parquet 路径(文件或目录),或 HF Hub 数据集名。
|
||
层 2 只用题面(prompt-only),不读数据自带的任何 completion。"""
|
||
|
||
output_dir: str
|
||
"""checkpoint 与日志输出目录(远程必须落在 /data/zym 下)。"""
|
||
|
||
# ---- teacher(层 2 的核心旋钮)----
|
||
teacher_model: str = "Qwen/Qwen3-4B"
|
||
"""本地 HF teacher 模型名。非机器路径(HF Hub 名各机可复现),故给默认值。
|
||
非显然约束:必须与 student **同 tokenizer**——KL 是逐词表位对齐求和,词表不
|
||
一致则第 v 个分量对不上、相除无意义(docs/03 §1)。此约束在 U4 构造 Trainer 时
|
||
比对 get_vocab() 显式校验,不匹配即报错,不静默。"""
|
||
|
||
# ---- 数据(复用层 1 的抽取逻辑,同 seed 同子集)----
|
||
dataset_split: str = "train"
|
||
|
||
subset_size: int | None = 1000
|
||
"""随机抽取的子集大小;None = 全量。非显然约束:与层 1 同 seed 同 size 才能
|
||
在同一批题上对比 SFT 与蒸馏,否则两层看的是不同题、曲线不可比。"""
|
||
|
||
# ---- 序列预算(prompt 截断 + 生成上限,见类 docstring 为何无 max_length)----
|
||
max_prompt_length: int = 1024
|
||
"""prompt 单独预算(prompt-only collator 按此左截断)。"""
|
||
|
||
max_new_tokens: int = 1024
|
||
"""student on-policy 生成的 token 上限。与 max_prompt_length 之和即序列总长 T,
|
||
显存账(§5)按 T=2048 估算。"""
|
||
|
||
enable_thinking: bool = False
|
||
"""Qwen3 chat 模板思考开关,喂给 student 生成。非显然约束:与层 1 取同值,
|
||
否则 prompt 渲染文本变、生成分布与 SFT 基线不可比(docs/02 §2.3 边界契约)。"""
|
||
|
||
# ---- 蒸馏损失(式(2) 与 docs/03 §2.3 三副面孔)----
|
||
beta: float = 1.0
|
||
"""KL 方向系数。0=前向 KL(π_T‖π_θ, mode-covering),1=反向 KL(π_θ‖π_T,
|
||
mode-seeking)=**式(2)**,(0,1)=JSD 插值。默认 1 即论文式(2);参数保留是因为
|
||
前向/反向/JSD 是同一公式(U2 顺手覆盖),且层 5 的 KL 锚要用前向。"""
|
||
|
||
kl_temperature: float = 1.0
|
||
"""散度内 softmax 前除进两侧 logits 的温度(docs/03 §2.3)。升温放大尾部
|
||
"暗知识"排序。非显然约束:它与下面的 gen_temperature 是**两个不同**的温度
|
||
——这个调的是 loss 里分布的软硬,那个调的是采样的随机性;恰好都默认 1.0,
|
||
但改一个不影响另一个。默认 1.0 即式(2)(不做温度缩放)。"""
|
||
|
||
# ---- on-policy 生成采样(式(2) 的 y~π_θ 期望)----
|
||
gen_temperature: float = 1.0
|
||
gen_top_p: float = 1.0
|
||
"""student 生成轨迹的采样参数。默认 temperature=1.0/top_p=1.0 = 纯采样自 π_θ,
|
||
最忠实于式(2) 的 on-policy 期望(docs/03 §3 抉择原则:本质忠于论文)。
|
||
调低是拿保真度换"少生成垃圾",卡了再动。"""
|
||
|
||
# ---- 优化 ----
|
||
# 差异标注:层 1 SFT 用 2e-5(信号密集的真 token 监督);层 2 是蒸馏,从论文
|
||
# §5.1 的蒸馏 lr=1e-6。小 lr 在这里还有额外好处:式(2) 的反向 KL 会梯度爆炸
|
||
# (§4.1,on-policy 采到 teacher 眼中的烂 token 时 log(π_θ/π_T)→∞),小步长
|
||
# 帮训练在毛刺中存活——这毛刺本身是层 2 要观察的教学目标(docs/03 §6.3)。
|
||
learning_rate: float = 1e-6
|
||
|
||
per_device_train_batch_size: int = 4
|
||
"""非显然约束:白盒蒸馏的显存大头是 **两份**全词表 logits(student+teacher,
|
||
(B,T,V) bf16 各 ~2.5G@B=4/T=2048)+ log_softmax 中间量,比层 1 更紧。B=4 是
|
||
§5 估算值(student 训练全套 ~10G + teacher 推理副本 ~9G + 两份 logits ~15G,
|
||
A800-80G 起步安全),但**必须**在首次远程冒烟用 nvidia-smi 实测确认,OOM 阶梯:
|
||
先降 B 到 2、仍不够再开 gradient_checkpointing。"""
|
||
|
||
gradient_accumulation_steps: int = 4
|
||
"""全局 batch = 4(per_device) × 4(卡) × 4(累积) = 64,与层 1 保持一致。"""
|
||
|
||
num_train_epochs: int = 1
|
||
|
||
max_steps: int = -1
|
||
""">0 时覆盖 num_train_epochs——远程 50 步冒烟用(docs/03 §6.3);-1 = 按 epoch。"""
|
||
|
||
lr_scheduler_type: str = "linear"
|
||
warmup_ratio: float = 0.0
|
||
|
||
max_grad_norm: float = 1.0
|
||
"""梯度裁剪阈值。此前是 HF Trainer 的静默默认(1.0),现显式化——它是式(2)
|
||
反向 KL 梯度爆炸(§4.1)的**隐形稳定器**:on-policy 采到 teacher 眼中烂 token
|
||
时单步梯度范数可炸到十几(2026-07-19 首冒烟实测 grad_norm 14→2),HF 默认
|
||
裁到 1.0 才让 loss 曲线平稳。把它设得远大于实测范数(≈关闭裁剪)可暴露原始
|
||
爆炸,供教学对照(train_whitebox.py 的 noclip 模式)。非显然约束:日志里的
|
||
grad_norm 是**裁剪前**范数,故 14→2 那串本身就是爆炸证据,只是被裁剪掩盖了。"""
|
||
|
||
gradient_checkpointing: bool = False
|
||
"""默认不开(§5 显存账 B=4 富余);OOM 时作为降 batch 之后的第二道降显存手段。
|
||
注意 FSDP 下此开关是 no-op(docs/02 §2.6),但层 2 坚持 DDP 故此处有效。"""
|
||
|
||
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:
|
||
"""构造即校验:配置错误必须在加载 4B teacher(GB 级下载)之前炸。"""
|
||
if not 0.0 <= self.beta <= 1.0:
|
||
raise ValueError(
|
||
f"beta 必须在 [0,1](0=前向/1=反向/中间=JSD),收到 {self.beta}"
|
||
)
|
||
if self.kl_temperature <= 0:
|
||
# 温度除进 logits,≤0 会翻转或炸掉分布
|
||
raise ValueError(f"kl_temperature 必须为正,收到 {self.kl_temperature}")
|
||
if self.gen_temperature <= 0:
|
||
# 非显然约束:0 在 HF 里是 greedy,会退化 on-policy 采样为确定性解码,
|
||
# 破坏式(2) 的 y~π_θ 期望;要纯 on-policy 就必须 >0
|
||
raise ValueError(
|
||
f"gen_temperature 必须为正(0=greedy 破坏 on-policy),"
|
||
f"收到 {self.gen_temperature}"
|
||
)
|
||
if not 0.0 < self.gen_top_p <= 1.0:
|
||
raise ValueError(f"gen_top_p 必须在 (0,1],收到 {self.gen_top_p}")
|
||
if self.max_new_tokens <= 0:
|
||
raise ValueError(f"max_new_tokens 必须为正,收到 {self.max_new_tokens}")
|
||
if self.learning_rate <= 0:
|
||
raise ValueError(f"learning_rate 必须为正,收到 {self.learning_rate}")
|
||
if self.max_grad_norm <= 0:
|
||
# 用远大于实测范数的值≈关闭裁剪;≤0 无意义(0 会把梯度裁没)
|
||
raise ValueError(f"max_grad_norm 必须为正,收到 {self.max_grad_norm}")
|
||
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}"
|
||
)
|