用 AI 核对知识蒸馏:损失乘了温度平方,梯度为何仍没完全变回去

昨天 3阅读

把较大的分类模型压缩到小模型时,知识蒸馏会让学生学习教师给出的整组概率。相关实现常把软目标损失乘T²。若AI把这解释成“温度影响被精确抵消”,一个三分类算例就能发现不对:补偿后的梯度仍可能变化。

用 AI 核对知识蒸馏:损失乘了温度平方,梯度为何仍没完全变回去

AI模型生成的概念插图:教师分布经过温度设置传给学生;柱形仅作示意,不是实际概率、模型大小比例或软件截图。

教师与学生使用同一温度

固定教师logits为[ln(9),0,0],学生logits为[0,0,0]。对每组logits先除以正温度T,再做softmax。T=1时,教师目标q为[9/11,1/11,1/11];T=2时,指数第一项从9变成3,目标变为[3/5,1/5,1/5]。

学生三个logits相等,因此两种温度下的概率p都为[1/3,1/3,1/3]。这里教师被冻结,只算软目标交叉熵L=-Σqᵢln pᵢ,不加入硬标签损失、类别权重或批量平均。教师与学生的类别顺序必须一致,否则算式正确也会教错目标。

一次除以温度,还有一次分布变化

对学生原始logit zᵢ求导,得到(pᵢ-qᵢ)/T。分母的T来自logit缩放的链式求导;而p与q本身也随T变化,不能只把温度前的概率差原样抄过来。

T=1时,三项梯度为[-16/33,8/33,8/33],约[-0.484848,0.242424,0.242424]。T=2时,先算概率差[-4/15,2/15,2/15],再除以2,得到[-2/15,1/15,1/15],约[-0.133333,0.066667,0.066667]。

现在对T=2的软目标损失乘T²=4,梯度也整体乘4,变成[-8/15,4/15,4/15],约[-0.533333,0.266667,0.266667]。第一项绝对值是T=1时的1.1倍,并没有精确恢复为原来的0.484848。

平方补偿对应什么近似

蒸馏原论文讨论的高温近似中,softmax差异会进一步按约1/T缩小,再结合链式求导的1/T,形成约1/T²的梯度尺度。乘T²有助于在调整温度时保持软目标与硬目标的相对影响,但这不是任意温度、任意logits下逐元素恒等的公式。

本次程序还固定每个温度的教师目标,对学生logits做中心有限差分,分别核对三项导数。若在扰动学生参数时又重新改变教师目标,检查的就不再是本文的冻结教师损失。三项梯度之和为0,也可作为softmax平移不变性的辅助核对。

使用KL散度时也要辨认方向。若计算KL(q教师‖p学生),与这里的交叉熵只差一个对学生恒定的教师熵,因此学生梯度相同;把方向反过来则不能直接沿用这个结论。工具接口接收概率还是对数概率,也应按文档核对。

这份练习适合审查蒸馏损失的温度、缩放与归约实现。温度如何选、学生能否保留教师能力,以及与硬标签损失怎样配比,都需要真实任务验证。软化目标包含教师的知识,也可能保留其错误;计算补偿系数并不构成性能保证。

资料核对日期:2026年10月2日。算例为原创教学设定,已用独立程序复核,未进行真实模型训练或效果测试。

参考资料

Hinton等:Distilling the Knowledge in a Neural Network,第2节

文章版权声明:除非注明,否则均为云鹊BLOG原创文章,转载或复制请以超链接形式并注明出处。