知识蒸馏的核心目标,是让一个轻量级学生网络去拟合教师网络的输出分布。真正动手训练时,学生模型精度可能反而不如单独用硬标签训练。这个现象通常不是蒸馏思路有问题,而是损失函数和温度参数没有配合好。蒸馏损失不是简单地把教师输出当成标注,它需要处理尖锐分布、梯度尺度、错误软标签等问题。

一、蒸馏精度损失到底从哪里来
教师网络输出的 logits 经过 softmax 后,在分类任务中往往非常尖锐:正确类别概率接近 0.99,其他类别接近 0。这种分布能提供给学生的额外信息很少,因为类别间相似度被压缩。解决思路是引入温度 T,对 logits 做软化:softmax(logits / T)。T 越大,输出越平滑,类别间的关系越明显。可是如果只做软化却不调整损失尺度,训练信号会被大幅减弱。
这里有一个容易忽略的点:软化后的梯度会缩小。假设教师输出分布不变,直接计算 KL 散度时,T 越大,损失值越小,梯度也越小。如果不做尺度补偿,软标签损失对总损失的贡献会被硬标签交叉熵淹没,学生网络等于几乎没有从教师那里学到类别间的相似结构。另一个原因是学生容量不足,无法完全复现教师在高维空间中的决策边界,过度依赖软标签又可能把教师错误预测放大。
可以用一个图像分类场景来观察:教师模型为 ResNet18,学生模型为 MobileNetV2,数据集为 CIFAR-100。温度 T 设为 1 时,教师输出接近 one-hot,蒸馏收益很小;T 设为 4 时,类别间相似关系被释放,学生验证精度通常明显提升;T 继续提高到 20,分布过度平滑,判别性反而下降。由此可见,精度损失一方面来自软目标信息不足,另一方面来自温度与损失权重的失衡。
二、损失函数设计:软标签与硬标签如何协同
标准蒸馏损失由两部分组成:硬标签交叉熵和软标签 KL 散度。常见写法为:L = alpha * L_CE(student_logits, y) + (1 - alpha) * T^2 * KL(log_softmax(student_logits / T), softmax(teacher_logits / T))。这里乘 T² 不是可有可无的系数,它来源于 softmax 泰勒展开后的梯度缩放。如果不乘,温度越高软损失越小,学生几乎学不到教师分布,蒸馏会退化为普通分类训练。
硬标签交叉熵提供无偏监督信号,防止教师模型在个别样本上给出错误高置信度时误导学生。alpha 通常设 0.5 到 0.9,教师越强、越可信,alpha 可以越小,也就是软标签权重大。但 alpha 太小时,学生可能过度模仿教师的噪声,尤其是小数据集上教师过拟合的样本。实际调参时,很多工程师习惯先固定 alpha 为 0.9,把温度调好,再反过来微调 alpha,这样收敛更快。
损失函数也不是只有 KL 散度一种选择。MSE loss 可以直接对齐学生与教师的软化 logits,但它对 logits 量级更敏感,通常需要先归一化或改成对齐概率。特征蒸馏、注意力图蒸馏可以在中间层补充监督,但工程落地时建议先调好输出层损失组合,再考虑中间层约束。输出层组合已经能解决大多数精度损失问题,中间层蒸馏更多是为了提升泛化性。
三、温度参数的作用与调参实践
温度 T 决定教师分布的平滑程度。T=1 时几乎无软化;T 在 3 到 10 之间,教师输出会呈现明显层次,类别间关系开始被学生感知。T 过高时,所有类别概率趋于均匀分布,KL 损失虽然乘了 T²,但类别间结构被抹平,模型可能只学到背景类的大致分布,反而降低判别性。T 过低则恢复尖锐,蒸馏收益很小。
实践中图像分类常用 T=4,NLP 任务如从 BERT 蒸馏到轻量编码器时,T=2 或 3 更常见。学生与教师容量差距较大时,可适当提高 T 到 6 到 10,因为学生更需要类别间关系来弥补表达能力的不足。建议先做一个 T 网格实验:固定 alpha=0.9,T 取 1、2、4、6、8、10,在小规模验证集上观察精度。需要注意 T 变化时 alpha 也要联合微调,不能只动一个参数。
动态温度策略也有人在用:训练初期高 T 让学生学习更宽泛的关系,后期降低 T 强化正确类别。但动态温度引入额外超参数,未必比固定 T 更好。更稳妥的做法是先把教师输出缓存下来,对软化分布做置信度校准,再选择固定 T。这样既节省显存,也让温度选择有依据,避免盲目调参。
四、PyTorch 实现示例与训练建议
下面是一个可直接使用的蒸馏损失函数实现。学生侧使用 log 概率,教师侧使用概率,KL 散度乘以 T² 补偿梯度缩放。
import torch
import torch.nn as nn
import torch.nn.functional as F
def distillation_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.9):
# 学生侧使用 log 概率,教师侧使用概率
soft_student = F.log_softmax(student_logits / T, dim=1)
soft_teacher = F.softmax(teacher_logits / T, dim=1)
# KL 散度,乘以 T^2 补偿软化带来的梯度缩放
distill_loss = F.kl_div(soft_student, soft_teacher, reduction='batchmean') * (T * T)
# 硬标签交叉熵
hard_loss = F.cross_entropy(student_logits, labels)
loss = alpha * hard_loss + (1.0 - alpha) * distill_loss
return loss
训练循环中,教师模型必须保持 eval 模式,并用 torch.no_grad() 包住教师推理,既省显存又避免给教师传梯度。以下是一个典型的小批量训练片段。
teacher.eval()
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
student_logits = student(images)
with torch.no_grad():
teacher_logits = teacher(images)
loss = distillation_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.9)
loss.backward()
optimizer.step()
if step % 100 == 0:
print(f"step {step}, loss {loss.item():.4f}\n")
训练时建议记录软损失和硬损失分别的数值,不要只看总损失。如果软损失一直很小,说明温度过低或软目标权重过低,需要重新检查 T 和 alpha。总损失下降但验证精度不涨甚至下降时,优先降低 T 或增大硬标签权重,而不是继续增加训练轮数。知识蒸馏的调参并不复杂,关键是把损失尺度、温度、教师输出质量三件事同时考虑清楚。