pytorch知识蒸馏核心是正确使用nn.kldivloss:student输出需f.log_softmax(logits/t),teacher输出需f.softmax(logits/t),温度t一致且通常取3–7,loss乘t²并设reduction='batchmean',teacher.eval()且用torch.no_grad(),student保持train()模式。

PyTorch 中知识蒸馏的核心不是换模型,而是改 loss —— 用 nn.KLDivLoss 替代交叉熵,并确保 student 输出经 F.log_softmax、teacher 输出经 F.softmax(且温度一致)。
为什么 KLDivLoss 必须配 log_softmax 和 softmax?
因为 nn.KLDivLoss 的数学定义是:KL(p||q) = Σ p·log(p/q),它要求第一个输入是 log-probabilities(即对数概率),第二个是 probabilities(即原始概率)。直接喂 raw logits 会得到错误梯度甚至 NaN。
- student 分支必须用
F.log_softmax(logits_student / T, dim=1) - teacher 分支必须用
F.softmax(logits_teacher / T, dim=1)(不能用log_softmax!) - 温度
T必须在两边严格一致,否则 KL 散度失去意义 - 别忘了
reduction='batchmean'—— 默认'mean'会对每个类求均值,导致梯度缩放异常
如何组织训练循环:teacher 不更新、student 只反向传播一次
常见错误是把 teacher 放进 optimizer.step() 或开了 requires_grad=True 却没 no_grad() 包裹。teacher 只能推理,student 才参与梯度计算。
- teacher 模型提前设为
eval(),且所有调用必须包在with torch.no_grad():内 - student 模型保持
train(),loss 计算后只对 student 参数调用optimizer.step() - 不要复用同一 batch 的 label 去算 teacher logits —— 如果 teacher 输入和 student 不同(如不同分辨率、预处理),必须各自前向
- 若用混合精度(
amp),确保torch.no_grad()在autocast外层,否则可能触发非预期的梯度计算
怎么加权蒸馏 loss 和原始任务 loss?
蒸馏效果差,往往不是模型问题,而是 alpha 和 T 没调好。典型组合是 alpha=0.7, T=3,但实际取决于 teacher/student 容量比。
- 总 loss =
alpha * kd_loss + (1 - alpha) * ce_loss,其中ce_loss是 student 对真实 label 的nn.CrossEntropyLoss -
kd_loss必须用nn.KLDivLoss(reduction='batchmean'),不能用nn.MSELoss(那是响应蒸馏,不是知识蒸馏) - 温度
T过高(如 >20)会让 softmax 输出过于平滑,teacher 知识“稀释”;过低(如 ≤1)则接近硬标签,失去软标签优势 - 建议用
torch.cuda.amp.GradScaler配合,因 KL loss 在低精度下更容易溢出
真正难的不是写完蒸馏代码,而是让 student 在小规模数据上稳定收敛 —— 这时 teacher 的 logits 是否干净、batch size 是否足够支撑 soft target 统计特性、以及 early stopping 是否基于蒸馏 loss 而非 val acc,都比模型结构本身更关键。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











