知识蒸馏论文入门指南:从模型压缩到落地实践

1次阅读
没有评论

共计 1370 个字符,预计需要花费 4 分钟才能阅读完成。

image.webp

背景:为什么需要知识蒸馏?

知识蒸馏(Knowledge Distillation)最早由 Hinton 在 2015 年提出,核心思想是让小型学生模型(Student)模仿大型教师模型(Teacher)的「思考方式」。这种技术能实现:

知识蒸馏论文入门指南:从模型压缩到落地实践

  • 模型压缩 :将 BERT 等大模型压缩到手机端
  • 知识迁移 :用 ImageNet 预训练模型指导医疗影像小模型
  • 效果提升 :学生模型甚至可能超越教师模型(神奇吧?)

三篇经典论文拆解

1. Hinton 开山之作(2015)

核心创新点就两个:

  1. Soft Target:不再用硬标签(0/1),而是用教师模型输出的概率分布(比如猫:0.8,狗:0.15,狐狸:0.05)
  2. 温度参数 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

五大避坑指南

  1. 温度参数 T 怎么选
  2. 分类任务通常 T =3~20
  3. 试试网格搜索:画个 T 和准确率的曲线图

  4. 学生模型太小怎么办

  5. 先用正常训练「暖机」
  6. 逐步增加蒸馏损失权重

  7. 教师模型过强反而不利

  8. 解决方案:用模型集成(多个教师投票)

  9. 中间层对齐的维度问题

  10. 加个 1 ×1 卷积调整通道数

  11. 蒸馏后模型变慢

  12. 检查是否有不必要的计算图保留

还能玩出什么花样?

  • 蒸馏 + 量化 :先蒸馏再量化,模型直接瘦身 90%
  • 蒸馏 + 剪枝 :让教师指导剪枝后的学生
  • 自蒸馏 :模型自己教自己(没错,真可以)

延展阅读

  1. 《Distilling Task-Specific Knowledge from BERT》
  2. 《TinyBERT: Distilling BERT for Natural Language Understanding》
  3. 《Self-Distillation: Towards Efficient and Compact Neural Networks》

最后说句大实话:看完这篇就去跑代码吧,论文里的数学公式看不懂?没关系,先让代码跑起来再说!

正文完
 0
评论(没有评论)