2.1 答案里的分寸感:基于Logits的响应蒸馏


文档摘要

2.1 答案里的分寸感:基于Logits的响应蒸馏 本节摘要:响应蒸馏是让学生模仿教师最终输出(Logits 或概率分布)的教法,是 Hinton 原始论文的标准形态,也是工程上最常用的第一选择。本节讲清它教的是什么知识、损失怎么写、为什么简单却不低端,并给出一套可直接跑的 PyTorch 实现与完整案例。 第一章铺垫了软目标的概念,本章开篇先讲最靠近答案的教材。响应蒸馏处于全册知识形态序列的第一站,下一节的特征蒸馏将往教师内部多走一层。 老师傅只给结论,能教出真本事吗 行话里管教师最后一层的原始打分叫 Logits——softmax 之前、不带归一化的那组数。响应蒸馏让学生去对齐的就是它。听起来徒弟只学了"结论",似乎肤浅,但 1.

2.1 答案里的分寸感:基于Logits的响应蒸馏

本节摘要:响应蒸馏是让学生模仿教师最终输出(Logits 或概率分布)的教法,是 Hinton 原始论文的标准形态,也是工程上最常用的第一选择。本节讲清它教的是什么知识、损失怎么写、为什么简单却不低端,并给出一套可直接跑的 PyTorch 实现与完整案例。

第一章铺垫了软目标的概念,本章开篇先讲最靠近答案的教材。响应蒸馏处于全册知识形态序列的第一站,下一节的特征蒸馏将往教师内部多走一层。

老师傅只给结论,能教出真本事吗

行话里管教师最后一层的原始打分叫 Logits——softmax 之前、不带归一化的那组数。响应蒸馏让学生去对齐的就是它。听起来徒弟只学了"结论",似乎肤浅,但 1.2 已经演示过:结论是一份完整分布,分布里藏着教师对每个类别相对关系的判断。学生模仿这份分布,等于同时学了三样东西:哪个类别对(硬标签也教)、哪些类别容易混(硬标签不教)、混到什么程度(只有分布教)。

一个常被低估的事实是:响应蒸馏在大量任务上是性价比最高的教法。它不窥探教师内部、不挑结构、实现只要十几行,而收益常达到更复杂教法的大半。工程上的合理顺序是先把响应蒸馏做扎实,确认收益边界后再考虑是否要更深的教材。

图 2-1 响应蒸馏的数据流:两次前向,一次更新

图 2-1 响应蒸馏的数据流:两次前向,一次更新

损失函数与实现

标准形态下,学生的总损失由两部分构成:对真实标签的交叉熵(硬损失),与对教师软化分布的 KL 散度(软损失)。KL 写成让学生分布逼近教师分布的方向:

import torch import torch.nn as nn import torch.nn.functional as F def response_kd_loss(student, teacher, x, y, T=4.0, alpha=0.9): """一轮响应蒸馏的损失。alpha 是软损失权重,常见起点 0.5-0.9。""" s_logits = student(x) # 学生原始打分 with torch.no_grad(): t_logits = teacher(x) # 教师打分,不回传 # 软损失:高温分布的 KL,乘 T 平方补偿被缩小的梯度 soft = F.kl_div( F.log_softmax(s_logits / T, dim=-1), F.softmax(t_logits / T, dim=-1), reduction="batchmean", ) * T * T # 硬损失:常温下对真实标签 hard = F.cross_entropy(s_logits, y) return alpha * soft + (1 - alpha) * hard # 输出说明:返回标量损失,调用 loss.backward() 后优化器只 # 会更新 student 的参数——teacher 的参数被 no_grad 隔离, # 且建议 teacher.eval() 关掉 dropout 等随机性,保证软目标稳定。

两个实现细节决定成败。其一,teacher.eval()no_grad 缺一不可:教师若带着 dropout 打分,软目标每轮抖动,学生学到的分寸感是模糊的。其二,软化用同一温度处理师生两侧的 logits——1.2 讲过除以温度后梯度缩小 T 平方倍,* T * T 就是补偿项。漏掉它时,alpha 需要人为放大几十倍才能让软损失起作用,很多"蒸馏没效果"的事故根源在这里。

一条完整案例:把响应蒸馏用到 CIFAR-10 级别任务

背景。 团队要为手机端相册做一个十类场景识别(海滩、夜景、文档、食物等),教师是服务端 ResNet 级模型,测试集 95.1%;端上预算要求学生参数量小于 300 万、单帧 15 毫秒。

操作。 学生选一个窄通道的同族网络。分三组实验:A 组学生纯硬标签训练;B 组加响应蒸馏(温度 4,alpha 0.9,即软损失占九成);C 组蒸馏并叠加标准数据增强。教师输出采用离线打分预存,训练机只驮学生。全程记录测试集准确率与端上延迟。

# 离线打分后的训练循环骨架:训练机不再加载教师 import torch def train_epoch(student, loader, soft_labels, optimizer, T=4.0, alpha=0.9): student.train() for i, (x, y) in enumerate(loader): s_logits = student(x) t_soft = soft_labels[i * x.size(0):(i + 1) * x.size(0)].to(x.device) hard = F.cross_entropy(s_logits, y) soft = -(t_soft * F.log_softmax(s_logits / T, dim=-1)).sum(-1).mean() soft = soft * T * T loss = (1 - alpha) * hard + alpha * soft optimizer.zero_grad() loss.backward() optimizer.step() # 注意:此时软损失用"教师概率 x 学生对数概率"的交叉熵形式, # 与 KL 形式在教师分布固定时等价,但省掉了再次前向教师。 # 输出说明:循环里只更新 student;soft_labels 是预存的 N x C 软化概率张量。

结果。 A 组 89.4%;B 组 92.6%,比 A 高 3.2 个点;C 组 93.0%。学生 int8 量化后 11 MB,端上单帧 9 毫秒,预算内达标。

解读。 三个观察值得记。第一,alpha 取 0.9(软损失为主)在这类分类任务上是常见甜点位:标签的信息量小(每样本一位 one-hot),教师分布的信息量大,重心自然偏软。第二,蒸馏与数据增强是叠加关系不是替代关系——C 组显示增强的收益没有被蒸馏吃掉。第三,离线打分让训练显存砍半,代价是增强只能用确定性变换;若训练时做随机增强,就得回到在线打分让教师实时批改。

变式。 若教师与学生类别数不一致(教师是 1000 类预训练模型,学生只认 10 类),可以在教师 logits 上做类目映射或只取相关子集;若教师对部分训练样本置信度极低,按 1.3 的经验过滤后再蒸;想进一步压延迟,把本节方案与 4.6 的量化感知蒸馏串联。

⚠️ 常见坑:把软损失权重调到 1.0、完全丢掉硬标签。教师在训练集上并非全对,硬标签是唯一的"外部纠错源";除非做纯无标签蒸馏,alpha 留一点给硬损失通常更稳。

💡 关键直觉:响应蒸馏的收益主要来自"难例附近的分寸感"。做消融时只统计整体准确率会低估它——把错误样本单独拉出来看,学生会发现自己在易混类目对(猫狗、沙滩湖畔)上的错误率下降最明显。

本节要点回顾

  • 教的什么:教师最终输出的完整分布——对错、易混类目、混淆程度,一次教齐。
  • 损失构成:硬交叉熵加软 KL,软损失乘温度平方补偿,alpha 常见起点 0.5 到 0.9。
  • 两个实现铁律:教师 eval 加 no_grad;师生同一温度软化。
  • 离线打分:教师软化输出预存后,训练显存减半、循环更简单,代价是增强受限。
  • 性价比定位:十几行代码拿到更复杂教法的大半收益,永远是蒸馏的第一站。

作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U