第 1 章 知识蒸馏原理详解 本章是全教程的理论核心。我们将从 Hinton 2015 的奠基论文出发,彻底搞懂:温度为什么能"软化"分布?什么是"暗知识"?为什么损失要乘 T²?KL 散度在蒸馏里到底度量了什么? 如果你只读一章,就读这一章。 1.1 从一个直觉说起:硬标签 vs 软标签 先看一个语言模型的经典场景:预测 "the cat" 的空格。 硬标签(Hard Label) 普通训练中,每个样本只有一个「正确答案」: 模型只被告知「lazy 是对的」,其他所有词都被同等视为「错」。这种信息量很稀薄——它不区分 "quick"(也合理)和 " Refrigerator"(完全无关)。
本章是全教程的理论核心。我们将从 Hinton 2015 的奠基论文出发,彻底搞懂:温度为什么能"软化"分布?什么是"暗知识"?为什么损失要乘 T²?KL 散度在蒸馏里到底度量了什么?
如果你只读一章,就读这一章。
先看一个语言模型的经典场景:预测 "the ___ cat" 的空格。
普通训练中,每个样本只有一个「正确答案」:
the [lazy] cat → 标签: lazy (词表中的某个 id)
模型只被告知「lazy 是对的」,其他所有词都被同等视为「错」。这种信息量很稀薄——它不区分 "quick"(也合理)和 " Refrigerator"(完全无关)。
一个训练好的教师模型,面对同样的输入,会输出一整条概率分布:
the [?] cat → 教师分布: lazy : 0.50 ← 最优 quick : 0.30 ← 次优,也合理 black : 0.10 ← 合理 hungry : 0.05 the : 0.02 ... 其余上万个词共分 0.03
这条分布携带的信息远比单一标签丰富。它告诉学生:
这种次优选项之间的相对关系,Hinton 称之为 暗知识(Dark Knowledge)——它们藏在 softmax 输出的"暗处",正常只看 argmax 时会被丢弃,但对学习极其宝贵。
核心洞察:蒸馏的本质,是让学生不仅学「正确答案」,还要学「教师认为哪些答案彼此相似、哪些不可能」。这种结构化的相似性信息,是教师在大规模语料上习得的"常识"。
问题来了:教师输出的原始 softmax 分布,往往过于尖锐(one-sided)——最优词概率接近 1,次优词概率小到几乎为 0,暗知识被"压扁"了。
普通 softmax (T=1): lazy:0.97 quick:0.02 black:0.005 ... ← 次优信息几乎看不见
温度(Temperature)T 就是用来解决这个问题的。修改后的 softmax 公式:
softmax_T(z)_i = exp(z_i / T) / Σ_j exp(z_j / T)
T = 1:标准 softmax,分布可能很尖锐。T > 1:分布变「软」(更平坦),次优选项的概率被放大,暗知识浮现。T → ∞:分布趋近均匀分布(所有词概率相等)。T < 1:分布更「硬」(更尖锐)。用代码直观感受温度的影响:
import torch import torch.nn.functional as F # 假设这是教师对 5 个词的原始 logits logits = torch.tensor([8.0, 6.0, 4.0, 2.0, 0.0]) for T in [1.0, 2.0, 4.0, 8.0]: probs = F.softmax(logits / T, dim=-1) print(f"T={T:<4} 分布: {[f'{p:.3f}' for p in probs.tolist()]}")
运行输出(示意):
T=1.0 分布: ['0.866', '0.117', '0.016', '0.002', '0.000'] ← 次优几乎消失 T=2.0 分布: ['0.624', '0.230', '0.085', '0.031', '0.011'] ← 次优浮现 T=4.0 分布: ['0.420', '0.268', '0.171', '0.109', '0.032'] ← 更平坦 T=8.0 分布: ['0.299', '0.248', '0.205', '0.155', '0.093'] ← 接近均匀
可以看到,T 越大,分布越平坦,次优选项的相对关系越清晰。这就是温度的作用——放大暗知识。
实践建议:蒸馏中通常取
T ∈ [2, 10],本教程默认T = 2。关于 T 的调参细节,我们会在《第 10 章 进阶实战》中详细讨论。
有了温度,我们就能定义蒸馏的总损失。它由两部分加权而成:
L_total = α · CE(student, 硬标签) + (1 - α) · T² · KL(teacher_soft ‖ student_soft) └─────────────────────┘ └────────────────────────────────────────┘ 硬标签项 软标签项(蒸馏项)
这部分和普通训练完全一样:让学生预测真实的下一个 token。
CE = -log P_student(真实token | 输入)
它保证学生至少能学会「正确答案」,不至于只顾着模仿教师而忘了真实标签。
这部分是蒸馏的灵魂:让学生在温度 T 下,模仿教师的软化分布。
teacher_soft = softmax_T(teacher_logits) student_soft = softmax_T(student_logits) KL项 = KL(teacher_soft ‖ student_soft) = Σ_i teacher_soft_i · log(teacher_soft_i / student_soft_i)
它衡量「学生分布偏离教师分布的程度」,越小说明学生学得越像。
α ∈ [0, 1] 控制二者的权衡:
| α 取值 | 含义 |
|---|---|
α = 1.0 |
退化为普通训练,无蒸馏(只有 CE) |
α = 0.0 |
完全依赖教师软标签(只有 KL) |
α = 0.5 |
两者并重(本教程默认) |
实践中 α 常取 0.3 ~ 0.7,让软标签发挥主要作用,硬标签作为兜底。
这是初学者最常困惑的点,值得专门解释。
当我们对 softmax_T(z) 关于 z 求导时,会发现梯度被额外除以了一个 T:
∂softmax_T(z)_i / ∂z_j ≈ (1/T) · [普通softmax在T=1时的梯度]
也就是说,温度软化后,logits 收到的梯度被压缩为原来的约 1/T。如果不补偿,T 越大,软标签项对 logits 的更新就越弱,蒸馏信号被"稀释"。
为了保持软标签项的梯度量级与硬标签项相当(这样两者才能公平竞争),Hinton 给软标签项乘上 T²:
补偿后梯度 ≈ T² · (1/T) = T ← 量级被还原
严格推导涉及一些数学,但结论很简单:乘 T² 是为了抵消温度对梯度的缩放,让软硬标签项的优化力度平衡。
工程结论:实现时,软标签 KL 计算完毕后必须乘以
T * T。这是 Hinton 原文的约定,几乎所有蒸馏实现都遵循。我们会在《第 6 章》的代码中看到这一行。
KL 散度(Kullback-Leibler divergence)是信息论中的概念,度量两个概率分布 P、Q 的差异:
KL(P ‖ Q) = Σ_i P_i · log(P_i / Q_i)
性质(很重要):
KL(P ‖ Q) ≥ 0,当且仅当 P = Q 时取 0。KL(P ‖ Q) ≠ KL(Q ‖ P),所以使用时要注意方向。蒸馏目标是「让学生逼近教师」,我们最小化的是:
KL(teacher_soft ‖ student_soft)
即把教师分布 P 当作"真值",衡量学生分布 Q 偏离它多远。注意 PyTorch 的 F.kl_div 函数签名是 kl_div(log_input, target),它内部计算的是 KL(target ‖ input),所以调用时第一个参数要传学生的 log_softmax,第二个传教师的 softmax。这个细节在《第 6 章》会专门强调,是新手最容易踩的坑。
import torch import torch.nn.functional as F # 教师分布(较平)与学生分布(较尖) teacher = F.softmax(torch.tensor([2.0, 1.0, 0.5, 0.0]), dim=-1) student_good = F.softmax(torch.tensor([2.0, 1.0, 0.5, 0.0]), dim=-1) # 完全一样 student_bad = F.softmax(torch.tensor([0.0, 0.5, 1.0, 2.0]), dim=-1) # 完全相反 kl_good = F.kl_div(student_good.log(), teacher, reduction='sum') kl_bad = F.kl_div(student_bad.log(), teacher, reduction='sum') print(f"KL(学生=教师): {kl_good:.4f} ← 接近 0,学得好") print(f"KL(学生相反): {kl_bad:.4f} ← 很大,学得差")
输出:第一个接近 0,第二个是一个明显的正数。这正是 KL 作为「蒸馏对齐指标」的价值——越小,学生越像教师。
你可能会问:学生模型这么小,直接用硬标签从零训练不就行了?为什么要多此一举搞蒸馏?
答案在于:软标签提供了更丰富的监督信号,相当于「数据增强 + 正则化」。
| 维度 | 硬标签 | 软标签 |
|---|---|---|
| 每个位置的信息量 | 1 个正确答案 | 整条概率分布(含上万个词的关系) |
| 对错误答案的区分 | 不区分(全标为"错") | 区分(次优 vs 极不可能) |
| 类似数据增强 | 否 | 是(隐式地告诉模型哪些词相似) |
| 正则化效果 | 弱 | 强(平滑分布抑制过拟合) |
工业界有大量证据表明蒸馏有效:
关于蒸馏效果的量化对比,我们会在《第 8 章 评估与对比》中给出具体的指标表格。
知识蒸馏不只有 Logits 蒸馏一种。下图展示了主要的蒸馏范式:
本教程聚焦最经典的 Logits 蒸馏(白盒):教师可达、对齐输出分布。它是入门的最佳起点,原理清晰、实现简洁、效果稳定。
其他范式的简要对比:
| 范式 | 迁移内容 | 代表方法 | 适用场景 |
|---|---|---|---|
| Logits 蒸馏 | 输出概率分布 | Hinton 2015、本教程 | 教师白盒、入门首选 |
| 中间层蒸馏 | 隐藏状态、注意力 | TinyBERT、MiniLM | 追求更高压缩比 |
| 序列级蒸馏 | 教师生成的文本 | SeqKD | 教师是黑盒 API |
| 关系蒸馏 | 样本间距离/角度 | RKD | 表征学习 |
中间层蒸馏和序列级蒸馏的进阶讨论,放在《第 10 章 进阶实战》。
记住这张图,它概括了蒸馏的全部要素:冻结的教师 + 可训练的学生 + 温度软化 + 软硬标签组合损失 + 仅更新学生。
| 概念 | 一句话理解 |
|---|---|
| 硬标签 | 单一正确答案,信息稀薄 |
| 软标签 | 教师输出整条概率分布,信息丰富 |
| 暗知识 | 软标签中次优选项的相对关系 |
| 温度 T | 软化分布,让暗知识浮现;T 越大越平坦 |
| α 权重 | 平衡硬标签 CE 与软标签 KL |
| T² 缩放 | 抵消温度对梯度的压缩 |
| KL 散度 | 度量学生分布偏离教师多远,越小越像 |
| 公式 | 含义 |
|---|---|
softmax_T(z)_i = exp(z_i/T) / Σ exp(z_j/T) |
带温度的 softmax |
L = α·CE + (1-α)·T²·KL(teacher‖student) |
蒸馏总损失 |
下一站:理论有了,接下来动手。在《第 2 章 环境准备》中,我们将搭好运行环境,并跑通第一次蒸馏训练。