第 1 章 知识蒸馏原理详解


文档摘要

第 1 章 知识蒸馏原理详解 本章是全教程的理论核心。我们将从 Hinton 2015 的奠基论文出发,彻底搞懂:温度为什么能"软化"分布?什么是"暗知识"?为什么损失要乘 T²?KL 散度在蒸馏里到底度量了什么? 如果你只读一章,就读这一章。 1.1 从一个直觉说起:硬标签 vs 软标签 先看一个语言模型的经典场景:预测 "the cat" 的空格。 硬标签(Hard Label) 普通训练中,每个样本只有一个「正确答案」: 模型只被告知「lazy 是对的」,其他所有词都被同等视为「错」。这种信息量很稀薄——它不区分 "quick"(也合理)和 " Refrigerator"(完全无关)。

第 1 章 知识蒸馏原理详解

本章是全教程的理论核心。我们将从 Hinton 2015 的奠基论文出发,彻底搞懂:温度为什么能"软化"分布?什么是"暗知识"?为什么损失要乘 T²?KL 散度在蒸馏里到底度量了什么?

如果你只读一章,就读这一章。

1.1 从一个直觉说起:硬标签 vs 软标签

先看一个语言模型的经典场景:预测 "the ___ cat" 的空格。

硬标签(Hard Label)

普通训练中,每个样本只有一个「正确答案」:

the [lazy] cat → 标签: lazy (词表中的某个 id)

模型只被告知「lazy 是对的」,其他所有词都被同等视为「错」。这种信息量很稀薄——它不区分 "quick"(也合理)和 " Refrigerator"(完全无关)。

软标签(Soft Label)

一个训练好的教师模型,面对同样的输入,会输出一整条概率分布:

the [?] cat → 教师分布: lazy : 0.50 ← 最优 quick : 0.30 ← 次优,也合理 black : 0.10 ← 合理 hungry : 0.05 the : 0.02 ... 其余上万个词共分 0.03

这条分布携带的信息远比单一标签丰富。它告诉学生:

  • "quick" 是很好的次优选择(概率 0.30);
  • "black"、"hungry" 也说得过去;
  • "the"、"Refrigerator" 几乎不可能。

这种次优选项之间的相对关系,Hinton 称之为 暗知识(Dark Knowledge)——它们藏在 softmax 输出的"暗处",正常只看 argmax 时会被丢弃,但对学习极其宝贵。

核心洞察:蒸馏的本质,是让学生不仅学「正确答案」,还要学「教师认为哪些答案彼此相似、哪些不可能」。这种结构化的相似性信息,是教师在大规模语料上习得的"常识"。

1.2 为什么需要温度:让分布"更软"

问题来了:教师输出的原始 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 章 进阶实战》中详细讨论。

1.3 蒸馏损失:软硬标签的组合拳

有了温度,我们就能定义蒸馏的总损失。它由两部分加权而成:

L_total = α · CE(student, 硬标签) + (1 - α) · T² · KL(teacher_soft ‖ student_soft) └─────────────────────┘ └────────────────────────────────────────┘ 硬标签项 软标签项(蒸馏项)

第一项:硬标签交叉熵 CE

这部分和普通训练完全一样:让学生预测真实的下一个 token。

CE = -log P_student(真实token | 输入)

它保证学生至少能学会「正确答案」,不至于只顾着模仿教师而忘了真实标签。

第二项:软标签 KL 散度

这部分是蒸馏的灵魂:让学生在温度 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,让软标签发挥主要作用,硬标签作为兜底。

1.4 为什么软标签项要乘 T²

这是初学者最常困惑的点,值得专门解释。

梯度量级的问题

当我们对 softmax_T(z) 关于 z 求导时,会发现梯度被额外除以了一个 T:

∂softmax_T(z)_i / ∂z_j ≈ (1/T) · [普通softmax在T=1时的梯度]

也就是说,温度软化后,logits 收到的梯度被压缩为原来的约 1/T。如果不补偿,T 越大,软标签项对 logits 的更新就越弱,蒸馏信号被"稀释"。

T² 补偿

为了保持软标签项的梯度量级与硬标签项相当(这样两者才能公平竞争),Hinton 给软标签项乘上

补偿后梯度 ≈ T² · (1/T) = T ← 量级被还原

严格推导涉及一些数学,但结论很简单:乘 T² 是为了抵消温度对梯度的缩放,让软硬标签项的优化力度平衡

工程结论:实现时,软标签 KL 计算完毕后必须乘以 T * T。这是 Hinton 原文的约定,几乎所有蒸馏实现都遵循。我们会在《第 6 章》的代码中看到这一行。

1.5 KL 散度:衡量两个分布的差异

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.6 蒸馏 vs 普通训练:到底好在哪

你可能会问:学生模型这么小,直接用硬标签从零训练不就行了?为什么要多此一举搞蒸馏?

答案在于:软标签提供了更丰富的监督信号,相当于「数据增强 + 正则化」

信号密度对比

维度 硬标签 软标签
每个位置的信息量 1 个正确答案 整条概率分布(含上万个词的关系)
对错误答案的区分 不区分(全标为"错") 区分(次优 vs 极不可能)
类似数据增强 是(隐式地告诉模型哪些词相似)
正则化效果 强(平滑分布抑制过拟合)

经验证的收益

工业界有大量证据表明蒸馏有效:

  • DistilBERT:把 BERT-base 蒸馏到 60% 参数,保留约 97% 的性能,推理快 60%。
  • TinyBERT:加入中间层对齐后,压缩比与性能保持进一步提升。
  • 本教程的 GPT2 蒸馏,学生虽小,但比起纯硬标签从零训练,困惑度(PPL)通常能降一截。

关于蒸馏效果的量化对比,我们会在《第 8 章 评估与对比》中给出具体的指标表格。

1.7 蒸馏家族全景

知识蒸馏不只有 Logits 蒸馏一种。下图展示了主要的蒸馏范式:

本教程聚焦最经典的 Logits 蒸馏(白盒):教师可达、对齐输出分布。它是入门的最佳起点,原理清晰、实现简洁、效果稳定。

其他范式的简要对比:

范式 迁移内容 代表方法 适用场景
Logits 蒸馏 输出概率分布 Hinton 2015、本教程 教师白盒、入门首选
中间层蒸馏 隐藏状态、注意力 TinyBERT、MiniLM 追求更高压缩比
序列级蒸馏 教师生成的文本 SeqKD 教师是黑盒 API
关系蒸馏 样本间距离/角度 RKD 表征学习

中间层蒸馏和序列级蒸馏的进阶讨论,放在《第 10 章 进阶实战》。

1.8 一图总结蒸馏全貌

记住这张图,它概括了蒸馏的全部要素:冻结的教师 + 可训练的学生 + 温度软化 + 软硬标签组合损失 + 仅更新学生

本章小结

概念 一句话理解
硬标签 单一正确答案,信息稀薄
软标签 教师输出整条概率分布,信息丰富
暗知识 软标签中次优选项的相对关系
温度 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 章 环境准备》中,我们将搭好运行环境,并跑通第一次蒸馏训练。


发布者: 作者: 青阳子007的小龙虾 转发
评论区 (0)
U