[LLM 7/10] 模型蒸馏:教师的错误答案,才是最有价值的部分
第 6 章我们把"上下文"蒸馏进了同一个模型的权重里,这一章我们要把"整个模型"蒸馏进一个更小的模型。 这两章是刻意成对的:context distillation 改变的是模型知道什么,让你不必再告诉它—— 模型蒸馏(model distillation)改变的是模型的大小,并尽力不改变它能做的事。 而本章的核心强烈违反直觉:教师能传给学生的最有价值的东西,不是正确答案, 而是教师犯错的方式——那个叫温度(temperature)的旋钮,正是让我们看见它的东西。
Open in Colab07_model_distillation.ipynb
1. 问题(Problem statement)
设想你走完了前六章,得到一个在组织的泰语任务上表现令人满意的 Qwen3-1.7B。 然后某一天,infrastructure 团队问了一个模型自己答不上来的问题:"GPU 一个月要花多少钱?"
1.7B 模型占用的显存几乎是 0.6B 的 3 倍,回答速度慢约 2.5 倍。 在每秒 100 个请求的量级上,这个差距不是细节——它是系统整个生命周期里每个月都要多买的显卡数量。
| 选项 | 质量 | 推理时的成本 |
|---|---|---|
| 直接部署 1.7B 教师 | 最好 | 显存约 3 倍、慢约 2.5 倍,为系统整个生命周期买单 |
| 直接部署 0.6B 学生 | 明显下降 | 便宜且快 |
| 用标准答案对学生做 SFT | 略有提升 | 便宜且快 |
| 模型蒸馏 | 向教师靠拢 | 便宜且快,和学生分毫不差 |
问题是:最后一行知道什么 SFT 那一行不知道的东西——同样是训练,同一份数据,差别到底在哪。
答案在于每个 token 携带的信息量。硬标签(hard label)是一个 one-hot 向量: 只说"答案是 3",就没了。而被温度软化后的教师分布说的是: "答案是 3——但 8 也不是完全没可能,至于'猫'就纯属胡扯了。" 正是这种在所有错误答案之上的排序,被 Hinton 称为暗知识(dark knowledge)。 它编码了世界的相似性结构(数字 3 离 8 比离猫更近), 而只要你只保留 argmax,它就消失得一干二净。
在每个 token 位置上,教师都有 151,936 个数字可给(Qwen3 的 vocab 大小)——硬标签只留下了其中一个。
2. 我们要做什么(Solution)
我们会拿 Qwen/Qwen3-1.7B 当教师、Qwen/Qwen3-0.6B-Base 当学生, 用三个层级来做蒸馏,教师传递的信息一级比一级精细:
- SeqKD——让教师生成回答,再用这些回答对学生做 SFT(来自分布的样本)
- Logit KD——让学生逐 token 位置模仿教师的整个分布(分布本身)
- GKD(可选加餐)——让学生采样自己的回答,教师在这些 token 上给出 on-policy(同策略)的分布评分
还有一个不可或缺的东西:对照行——用标准答案在同一份数据上、以相同的 step 数对学生做 SFT。 没有这一行,我们根本分不清收益究竟来自"教师的分布",还是仅仅来自"多训练了一会儿"。
硬标签说"答案是 3"——教师分布说的是"3,但 8 也接近对,而'猫'不可能"。 在错误答案之上的排序就是暗知识,而温度是揭示它的旋钮。 在 T = 1 时这份知识被压得几乎看不见(教师置信 0.97);把 T 提上去,它才变成真正可训练的 signal。
这不是实验室里的花招——世界上大多数"小而强"的模型正是这样被造出来的。 我们整个系列一直在用的 Qwen3-0.6B,本身就是用它家族里更大的型号做 strong-to-weak distillation 训练出来的。 这一章我们做的是同一件事,只是缩到免费 Colab 装得下的规模。
3. 公式(Equation)
3.1 Hinton 的 KD loss
- = 学生与教师在同一 input、同一 token 位置上的 logits
- = 真实标准答案(hard label)——第一项就是普通的 cross-entropy,和 SFT 一模一样
- = 温度,在 softmax 之前同时除进两侧的 logits
- = soft 项的权重(我们用 0.9——以听教师为主,让标准答案兜底防跑偏)
注意 KL 的方向:教师在前。这是 forward KL,它强迫学生把概率铺开, 覆盖教师给了权重的每一个地方。把这一点记住,待会儿公式 3.4 会让它变成"一条线上的一个点"。
那么乘在 KL 前面的 是从哪来的?互联网上几乎每一份 KD 代码都带着这个系数, 但解释原因的少之又少——而如果不理解它,你会在不知不觉中用错误的方式调 T。
3.2 推导 的来历——把"理解"和"照抄"分开的两行推导
第 1 行——soft 项对学生单个 logit 的梯度,就是 softmax-CE 的标准梯度, 再经过 的 chain rule 吐出一个 :
第 2 行——当 很大时,softmax 会在 uniform 附近被摊平: (其中 = vocab 大小), 因此差值 又按 再缩一层:
soft 项的梯度按 缩放,而 hard CE 项完全不依赖 。 如果不把 乘回去,把 T 从 1 调到 4 就等于偷偷把 soft 项的学习率除以约 16。 你会得出"T 调高了不管用"的结论,而实际上你只是不小心把自己的 soft loss 关掉了而已。 乘上 让梯度的尺度几乎不随 T 变化——调好的 在每个 T 下都保持原来的含义。
第 2 行还有一份赠品:在同一极限下,soft loss 退化成对中心化 logits 的 MSE 匹配—— KD 就是一种"软化版的 logit regression",给分布头部的权重高于尾部。
3.3 SeqKD——便宜的 baseline(Kim & Rush, 2016)
读起来眼熟吗?——这就是在教师生成的回答上做的普通 SFT,仅此而已。 教师没有把整个分布送过来,而是送来一个从自己的分布里采出的"单个样本"。
一个常被忽视的优点:SeqKD 完全不在乎 tokenizer 是否一致,因为它传的是文本,不是 logits。 很多宣传"distilled from GPT-4"的开源模型,实际上就是纯 SeqKD—— 通过 API 收集教师回答,然后 SFT。这正是它成为 baseline 的原因:在上更贵的方法之前,必须先把它测出来。
3.4 GKD 与 generalized JSD——把本章和第 6 章连起来的那条线
公式 3.1 的 logit KD 有一个结构性弱点:学生是在别人写的句子上学习的(teacher forcing), 但真正使用时它必须接着自己的回答往下生成——这种不断累积的偏差叫作 exposure bias。 GKD(Agarwal et al., 2023)的解法是让学生自己采样回答,教师在这些 token 上打分, 并顺手把分布之间的距离推广为:
其中 = 教师、 = 学生。 从一个极端扫到另一个极端:
- → forward KL ——mass-covering:学生必须铺开覆盖教师的每一个 mode
- → reverse KL(反向 KL)——mode-seeking(模式寻找):学生选择守住自己扛得动的那几个 mode
把话说到最明白:第 6 章选用的 reverse KL 并不是来自另一个世界的怪东西——它就是这条线上 的那个点,而 Hinton 的经典 KD 是 的那个点。 这两章因此是同一个家族的成员, 差别只在"谁必须向谁靠拢"——比教师小得多的学生往往在 mode-seeking 一侧受益更多, 因为它的容量本来就不足以覆盖教师的所有 mode。
这篇文章大约是本章的前 30%。其余部分——环境准备、数据准备、核心代码、实测结果与总结——都在免费的 LLM Finetuning 课程中,使用 Google 登录即可阅读。
在课程中阅读完整章节 →