HOMEWORK 06

自蒸馏(SDPO)

模型当自己的老师:用同组里做对的那条 rollout 当演示,把「看过答案的自己」的分布蒸回「没看过答案的自己」。整个循环在一张卡上就能看完。

对应章节:第 12 章 · 合成数据与蒸馏 参考实现:_src/code/distillation/ 配置:distillation/configs/sdpo.yaml 状态:⚠ 本机未实跑

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 的机制

一步训练是这样的:

  1. 学生 rollout:对一道 spell_backward 题目,只给题面,采 num_rollouts 条回复。
  2. 环境验证:每条由 Reasoning Gym 判定对错(得分 $\ge$ success_reward_threshold 即算正确)。
  3. 挑演示:从这一组里取一条正确的兄弟 rollout 当作演示(demonstration)。
  4. 自教师前向:把「题目 + 这条演示」拼成一个更长的 prompt,让同一份权重再前向一次。这就是「看过答案的自己」。
  5. 蒸馏:在学生采样出的那条序列的每个位置上,用 top-$K$ reverse KL 把自教师的下一 token 分布蒸回学生。

关键在第 3 步:如果这一组里一条正确的都没有,就没有演示可蒸,这个 prompt 被整个跳过。所以每一次更新都有一个真实的、来自模型自己的正确示范。

SDPO 和 GRPO 的关系 两者都在同一个任务(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_nameQwen/Qwen3-1.7B同时充当学生和自教师——同一份权重,两次前向,区别只在 prompt 里有没有演示。
data.specsspell_backwardReasoning Gym 任务混合。min_word_len/max_word_len 控制难度,也就控制了「组内有没有正确样本」。
num_rollouts8每题采几条。演示就是从这一组里选的,它直接决定 skipped 的高低。
success_reward_threshold1.0得分达到多少才算「可以当演示」。默认要求完全正确。
kl_top_k20蒸馏 KL 保留的 logits 数。实现会把剩余概率质量收进一个尾桶,使 top-$K$ 加尾桶构成合法的 $(K{+}1)$ 维分布。
prompts_per_step16每个优化步累积多少个 prompt 的梯度。
rollout_chunk4(在 config.py 里,YAML 未写)每次前向/反向处理几条 rollout。这是控制峰值显存的主要开关(第 3 节)。
max_prompt_len / max_reprompt_len512 / 1024学生 prompt 与教师 reprompt 的长度上限。教师那条更长,因为要塞进完整的演示。
lr / warmup_ratio1e-6 / 0.0恒定学习率、无 warmup。这个量级比 SFT 低得多,和 on-policy 蒸馏的稳定性经验一致。
num_steps200总优化步数。注意每一步都包含一轮生成,墙钟时间主要花在这里。
两条通道的 prompt 差在哪 学生看到的是题目;自教师看到的是题目 + 一条正确的兄弟解答(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 不可能产生多个答案块。
常见误区 「loss 在降,说明蒸馏在起作用。」不对——loss 只说明学生在向教师靠拢,不说明靠拢的方向是能力。学生完全可以通过模仿教师的风格(长度、句式、标签用法)把 KL 压下去,而解题能力一点没变。

判据只有一个: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 数」,分块梯度累加起来恰好等于整组梯度——数学上完全等价。
显存不够时按这个顺序调 1) 调小 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。调高梯度更稳但每步更慢,调低则更新更频繁但噪声更大。

一个看似矛盾的经验 上游针对更难的任务给出的稳定性经验正好相反:把 mini batch 拆小(每轮做 16 次优化步而不是 1 次)被列为「最大的杠杆」。

两者并不矛盾——调大 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 是否随训练下降?如果没有,你判断卡在哪一环?