自蒸馏(SDPO)
模型当自己的老师:用同组里做对的那条 rollout 当演示,把「看过答案的自己」的分布蒸回「没看过答案的自己」。整个循环在一张卡上就能看完。
0. 任务目标
homework/RESULTS.md:本组要在同一步里既生成又反传,且教师通道和学生通道各要一次完整前向,显存预算超出了本机剩余的约 6.9 GiB(详见第 3 节)。
因此,下面的内容是任务设计与预期观察点,不是实测结果。本页不给出任何具体的实验数值——只讲该盯哪些指标、每个指标健康时是什么形态、异常时说明什么。真正的实测记录只在
homework/RESULTS.md 里,那里目前只有 HW1、HW2、HW4 三组。
第 12 章讲了蒸馏的方向性问题:传统的知识蒸馏(word-level KD、sequence-level KD)都是off-policy 的——学生在教师生成的文本上学,而它自己推理时走的是自己的分布。这个错配就是曝光偏差(exposure bias):学生只在「教师走过的路」上被训练过,一旦自己走岔了,就没人教它怎么回来。
on-policy 蒸馏(on-policy distillation, OPD)的解法很直接:让学生自己生成,然后在学生自己走过的路径上问教师「你在这里会怎么选」。SDPO(Self-Distillation Policy Optimization)把这个思路推到极致——连教师都不要了,模型当自己的老师。
SDPO 的机制
一步训练是这样的:
- 学生 rollout:对一道
spell_backward题目,只给题面,采num_rollouts条回复。 - 环境验证:每条由 Reasoning Gym 判定对错(得分 $\ge$
success_reward_threshold即算正确)。 - 挑演示:从这一组里取一条正确的兄弟 rollout 当作演示(demonstration)。
- 自教师前向:把「题目 + 这条演示」拼成一个更长的 prompt,让同一份权重再前向一次。这就是「看过答案的自己」。
- 蒸馏:在学生采样出的那条序列的每个位置上,用 top-$K$ reverse KL 把自教师的下一 token 分布蒸回学生。
关键在第 3 步:如果这一组里一条正确的都没有,就没有演示可蒸,这个 prompt 被整个跳过。所以每一次更新都有一个真实的、来自模型自己的正确示范。
spell_backward)上、用同一种组采样(num_rollouts: 8)、靠同一个环境验证器。差别只在「把信号写回参数」的方式:
GRPO(HW3)用一个标量优势去加权整条序列的 log-prob —— 信号极稀疏,每条 rollout 只贡献一个数。 SDPO 用一个逐 token 的分布(top-$K$ 软标签)当监督 —— 信息密度高出好几个数量级。
代价是:SDPO 需要「组内至少有一条对的」才能启动,而 GRPO 只需要「组内有差异」。两者卡死的条件不同,但根源是同一个:模型得先偶尔能做对。
这一组作业的目标,就是把这个「先能对、再变强」的自举循环看清楚——尤其是它什么时候会卡死,以及卡死时三个指标各自会表现成什么样。
1. 运行命令与关键配置
cd _src/code/
uv sync
# 可选:把指标打到自己的 W&B
export WANDB_PROJECT=rlhf-book
python -m distillation.train --config distillation/configs/sdpo.yaml
同样建议用包装器把逐步指标落盘(原理见 HW4 第 6 节):
METRICS_JSONL=homework/hw6-distill/logs/metrics_sdpo.jsonl \
WANDB_MODE=disabled \
python homework/tools/run_with_metrics.py distillation.train \
--config distillation/configs/sdpo.yaml
sdpo.yaml 的关键旋钮
| 字段 | 默认 | 含义与调参影响 |
|---|---|---|
model_name | Qwen/Qwen3-1.7B | 同时充当学生和自教师——同一份权重,两次前向,区别只在 prompt 里有没有演示。 |
data.specs | spell_backward | Reasoning Gym 任务混合。min_word_len/max_word_len 控制难度,也就控制了「组内有没有正确样本」。 |
num_rollouts | 8 | 每题采几条。演示就是从这一组里选的,它直接决定 skipped 的高低。 |
success_reward_threshold | 1.0 | 得分达到多少才算「可以当演示」。默认要求完全正确。 |
kl_top_k | 20 | 蒸馏 KL 保留的 logits 数。实现会把剩余概率质量收进一个尾桶,使 top-$K$ 加尾桶构成合法的 $(K{+}1)$ 维分布。 |
prompts_per_step | 16 | 每个优化步累积多少个 prompt 的梯度。 |
rollout_chunk | 4(在 config.py 里,YAML 未写) | 每次前向/反向处理几条 rollout。这是控制峰值显存的主要开关(第 3 节)。 |
max_prompt_len / max_reprompt_len | 512 / 1024 | 学生 prompt 与教师 reprompt 的长度上限。教师那条更长,因为要塞进完整的演示。 |
lr / warmup_ratio | 1e-6 / 0.0 | 恒定学习率、无 warmup。这个量级比 SFT 低得多,和 on-policy 蒸馏的稳定性经验一致。 |
num_steps | 200 | 总优化步数。注意每一步都包含一轮生成,墙钟时间主要花在这里。 |
rollout.py::build_teacher_prompt)。但两者被要求预测的是同一条序列——也就是学生刚刚采样出来的那条。
所以蒸馏信号的含义是:「在你自己写出的这个位置上,如果你已经看过答案,你会怎么选下一个 token?」这就是 SDPO 里「特权上下文(privileged context)」的全部内容,也是它区别于普通自训练的地方。
2. 该盯的三个指标
train.py 每个优化步记录一组指标,其中 reward、loss、skipped 是必须一起读的三个(另有 grad_norm、lr、hours)。单看任何一个都会得出错误结论,这是本节的全部要点。
| 指标 | 是什么 | 健康的形态 | 不对劲时说明什么 |
|---|---|---|---|
reward |
本步所有 rollout 的平均环境得分 | 整体上行。这是学生能力真的在变强的唯一直接证据 | 躺平不动 = 学生没在学。先去看 skipped 是不是太高,再去看 loss |
loss |
掩码后的 top-$K$ reverse KL(学生 $\|$ 自教师) | 随着学生内化「演示条件下的分布」而下行 | 见下方专门讨论——它单独下降什么都不能证明 |
skipped |
polled - len(batches):轮询到的 prompt 里,整组 rollout 全错、没有演示可蒸的个数 |
训练初期较高,随学生变强单调下降 | 一直很高 = 任务对当前模型太难。轮询打到 prompts_per_step * 100 会直接抛异常,属于「快速失败」而不是静默卡死 |
为什么「loss 降但 reward 不动」是最危险的信号
这一条值得单独讲,因为它是本组作业最有教学价值的一个诊断。
蒸馏损失衡量的是学生的下一 token 分布和自教师有多接近。而自教师和学生的差别,理论上应该是「见过正确演示 vs 没见过」——所以拉近这个距离,应该等价于「让学生在没有演示时也表现得像见过演示一样」,也就是变得更会解题。
但这个等价关系有一个隐含前提:教师和学生的分布差异,主要来自「解题能力」而不是别的东西。如果演示挑得不好,这个前提就崩了:
- 如果总是挑最短的那条正确回复当演示,教师分布相对学生的主要差异就变成了「更短」。学生会忠实地学会把答案写短,而不是学会解题。上游文档明确记录了这个失败模式:在有多条正确回复时应当随机挑一条,挑最短的那条会让训练坍缩。
- 如果演示是「反复回溯、最后蒙对」的那种(生成里出现多个
<answer>块),学生学到的就是这种回溯行为本身。上游的建议是把这类样本过滤掉,更好的做法是直接把</answer>设成停止序列,从源头上让 rollout 不可能产生多个答案块。
判据只有一个:
reward 有没有跟着动。如果 loss 稳步下行而 reward 长期平坦,不要去调学习率——去检查演示是怎么挑的,并把教师 prompt 和学生 rollout 的原文拉出来逐条对比。文本会告诉你答案,曲线不会。
三者之间的动力学
这套方法的核心动力学是一个正反馈回路:
skipped 高 → 每步真正参与更新的 prompt 少 → 梯度噪声大 → reward 爬得慢 → skipped 继续高
这既是 SDPO 最容易卡死的地方,也解释了它的启动条件:模型必须「本来就偶尔能做对」,自举才能开始。一个完全不会做这道题的模型,在 SDPO 下会永远停在原地(并最终触发那个轮询上限的异常)——它没有任何可以蒸馏的正确示范。
反过来,训练后期 skipped 会自然降到很低,因为学生大部分时候都能做对。此时 reward 的增长也会放缓——可蒸的「信息增量」变少了。这时候该做的是提高任务难度,而不是继续跑更多步。
3. 显存:为什么本机跑不动
homework/RESULTS.md 给本组的记录是「教师与学生双模型驻留,同样受显存限制」。这个说法需要展开一下,因为它在一般的 on-policy 蒸馏和这份 SDPO 实现里含义不同——搞清楚这一点,你才知道该往哪里省显存。
| 场景 | 教师是什么 | 显存代价 |
|---|---|---|
| 一般 OPD(第 12 章讲的形态) | 一个独立的、通常更大的冻结模型 | 真的要驻留两份权重,教师那份还往往更大 |
| 本实现的 SDPO | 同一份权重,只是 prompt 里多了演示 | 权重只有一份,但每步要跑两条完整前向(学生一次带梯度、教师一次 no_grad),且两次都要物化稠密 logits |
所以这份实现真正的显存大头不是权重,而是 logits。loss.py 的注释说得很直白:学生的 logits [R, A, V] 及其梯度是主要开销——$R$ 是本组 rollout 数、$A$ 是 completion token 数、$V$ 是词表大小。词表动辄十几万,这个张量非常大。
实现用了两个技巧来压它:
logits_to_keep = A + 1:只把 completion 那一段过 LM head,prompt 段的 logits 根本不算。教师的 reprompt 很长(要塞演示),这一招省下的量相当可观。rollout_chunk分块反传:把 $R$ 条 rollout 切成小块,每块前向-反向完就释放,把峰值显存压到「一块」而不是「一组」。因为损失是「对 rollout 求和再除以全局 token 数」,分块梯度累加起来恰好等于整组梯度——数学上完全等价。
rollout_chunk(4 → 2 → 1)——这是零代价的,只换速度不换算法语义;
2) 调小 max_new_tokens 和 max_reprompt_len(直接减小 $A$);
3) 调小 kl_top_k——注意这会改变监督信号的质量(第 4 节);
4) 换更小的基座;
5) 最后才动 num_rollouts——它是学习信号的来源,砍它等于砍掉演示的供给,skipped 会立刻恶化。
另外一笔账:时间。上游 README 记录的参考运行是在一张 24 GB 消费级显卡上跑了不到 20 小时(这是上游的记录,不是本机的实测)。本组作业不要求跑完,跑到三条曲线的形态清楚可辨即可。
最后,本机还踩过一个和显存无关但同样致命的环境坑:上游为 x86_64 Linux 钉的是 cu126 轮子,其 get_arch_list() 不含 sm_120,而 RTX 5080 是 Blackwell。症状是 torch.cuda.is_available() 返回 True,但一启动 kernel 就报 no kernel image is available for execution on the device。修复方式见 homework/RESULTS.md。
4. 自选消融
复制 distillation/configs/sdpo.yaml,固定任务,只扫下面三个旋钮。它们分别对应本方法的三个关键维度:演示的供给、监督的密度、更新的稳定性。
消融一:num_rollouts(演示的供给)
默认 8,可以试 4 / 8 / 16。每个 prompt 采得越多,「组内至少有一条正确」的概率越高,skipped 直接下降。这是缓解第 2 节那个正反馈回路最直接的杠杆。
代价是每步生成量线性增加。所以这个消融真正要回答的是一个预算分配问题:同样的生成算力,是该多采几条(降低 skipped、每步更有效)还是多做几步(更新次数更多)?观察 reward 随累计生成量(而不是随步数)的曲线,才是公平的比较。
消融二:kl_top_k(监督的密度)
默认 20,可以试 5 / 20 / 50。它决定 reverse KL 匹配教师分布的宽度。
- 调高:监督更完整,学生看到教师分布更细的形状;代价是显存和算力。
- 调得太低:尾桶承担的概率质量过多,梯度信息变粗——极端情况下学生只能学到「排名第一的 token 是什么」,软标签的大部分价值就损失了,退化成接近硬标签的模仿。
这个消融的意义在于把「为什么蒸馏比 SFT 信息量大」这件事量化:$K=1$ 时几乎就是在做 SFT,$K$ 越大越接近完整的分布匹配。找出收益开始饱和的那个 $K$,就是这份任务上「软标签到底有多少信息」的答案。
消融三:prompts_per_step(更新的稳定性)
默认 16。调高梯度更稳但每步更慢,调低则更新更频繁但噪声更大。
两者并不矛盾——调大
prompts_per_step 是在降方差,拆小 mini batch 是在提高更新频率。哪个更重要取决于你离「奖励曲线能不能起来」有多远:曲线还没起来时,更新频率更值钱;曲线已经在爬、但抖得厉害时,降方差更值钱。
还值得一试的两个
- 降低
success_reward_threshold(从 1.0 降到某个部分分阈值)。这会让更多 rollout 有资格当演示,skipped下降;但演示的质量也跟着下降。观察reward的最终水平是不是也跟着降——这是一个非常干净的「演示质量 vs 演示数量」权衡实验。 - 改演示的挑选规则。默认是从正确的兄弟里选一条;试着改成「总是选最短的那条」,然后观察
loss和reward的背离。这是第 2 节那个「学风格不学能力」失败模式的主动复现——亲手把它做出来一次,以后在别的项目里就能一眼认出。
做任何一组消融时,都请把循环里打印的教师 / 学生 rollout 样本一起存下来。这套方法的很多问题(演示太短、反复回溯、格式漂移)在曲线上是看不出来的,只有文本会说话。
本组小结
| 问题 | 答案 |
|---|---|
| 三个核心指标? | reward(真信号)、loss(蒸馏是否在收敛)、skipped(有没有演示可学)。必须一起读。 |
| 最危险的信号? | loss 降但 reward 不动——学生在拟合教师的风格而非能力。去查演示挑选逻辑。 |
| 最容易卡死的地方? | skipped 高 → 有效 prompt 少 → 梯度噪声大 → reward 爬得慢 → skipped 继续高。正反馈回路。 |
| 缓解回路最直接的杠杆? | 调大 num_rollouts(组内至少一条对的概率上升),或降低任务难度。 |
| 本机为什么没跑? | 教师与学生两条通道的前向 + 稠密 logits 的显存开销,超出约 6.9 GiB 预算。 |
交付物清单
- 你实际使用的配置(若换了基座或改了任务难度,写清改动与理由)。
- 完整日志 + 逐步指标 JSONL,三条曲线:
reward、loss、skipped。 - 训练早期与后期各若干组「教师 prompt / 学生 rollout」原文对照——这比曲线更能说明学生到底在学什么。
- 一段说明:
skipped是否随训练下降?如果没有,你判断卡在哪一环?