共计 1370 个字符,预计需要花费 4 分钟才能阅读完成。
背景:为什么需要知识蒸馏?
知识蒸馏(Knowledge Distillation)最早由 Hinton 在 2015 年提出,核心思想是让小型学生模型(Student)模仿大型教师模型(Teacher)的「思考方式」。这种技术能实现:

- 模型压缩 :将 BERT 等大模型压缩到手机端
- 知识迁移 :用 ImageNet 预训练模型指导医疗影像小模型
- 效果提升 :学生模型甚至可能超越教师模型(神奇吧?)
三篇经典论文拆解
1. Hinton 开山之作(2015)
核心创新点就两个:
- Soft Target:不再用硬标签(0/1),而是用教师模型输出的概率分布(比如猫:0.8,狗:0.15,狐狸:0.05)
- 温度参数 T :通过调节 T 控制概率分布的「平滑度」,公式长这样:
q_i = \frac{exp(z_i/T)}{\sum_j exp(z_j/T)}
2. FitNets(2015)
发现光学输出不够,让学生直接模仿教师的中间层特征:
- 添加 Hint 层 :对齐教师和学生模型的中间层
- 适合场景:学生模型结构较深时
3. Attention Transfer(2017)
更聪明地利用中间层信息:
- 用注意力图(Attention Map)作为知识载体
- 代码实现只需加几行 CNN 的注意力计算
动手实现一个蒸馏器
用 PyTorch 写个最简单的图像分类蒸馏(完整代码见 GitHub):
# 教师模型(现成的 ResNet50)teacher = torchvision.models.resnet50(pretrained=True)
# 学生模型(自己搭的小网络)student = nn.Sequential(nn.Conv2d(3, 32, 3),
nn.ReLU(),
nn.Flatten(),
nn.Linear(32*30*30, 10) # 假设是 CIFAR-10
)
# 核心——蒸馏损失函数
def distillation_loss(student_logits, teacher_logits, T=3):
# 计算 KL 散度(记得加温度!)loss = F.kl_div(F.log_softmax(student_logits/T, dim=1),
F.softmax(teacher_logits/T, dim=1),
reduction='batchmean'
) * (T*T) # 乘以 T²保持梯度量级
return loss
五大避坑指南
- 温度参数 T 怎么选 :
- 分类任务通常 T =3~20
-
试试网格搜索:画个 T 和准确率的曲线图
-
学生模型太小怎么办 :
- 先用正常训练「暖机」
-
逐步增加蒸馏损失权重
-
教师模型过强反而不利 :
-
解决方案:用模型集成(多个教师投票)
-
中间层对齐的维度问题 :
-
加个 1 ×1 卷积调整通道数
-
蒸馏后模型变慢 :
- 检查是否有不必要的计算图保留
还能玩出什么花样?
- 蒸馏 + 量化 :先蒸馏再量化,模型直接瘦身 90%
- 蒸馏 + 剪枝 :让教师指导剪枝后的学生
- 自蒸馏 :模型自己教自己(没错,真可以)
延展阅读
- 《Distilling Task-Specific Knowledge from BERT》
- 《TinyBERT: Distilling BERT for Natural Language Understanding》
- 《Self-Distillation: Towards Efficient and Compact Neural Networks》
最后说句大实话:看完这篇就去跑代码吧,论文里的数学公式看不懂?没关系,先让代码跑起来再说!
正文完
