文章目录

知识蒸馏中教师-学生网络的KL散度与温度软化

发布于 2026-07-31 10:51:12 · 浏览 63 次 · 评论 0 条

知识蒸馏中教师-学生网络的KL散度与温度软化

知识蒸馏的核心,是让一个小模型(学生网络)去模仿一个大模型(教师网络)的“想法”,而不是只模仿最终的答案。关键在于,教师网络输出的是一个概率分布,里面包含了“这个类别像什么”的丰富信息。学生网络要学的,正是这种分布。而实现这种模仿的数学工具,是KL散度(Kullback-Leibler散度)。为了让这个分布“软化”到适合学习,还需要引入一个叫“温度”的参数。


1. 理解教师网络输出的“暗知识”

认识 教师网络的输出。一个分类模型,最后一般接一个 Softmax 层,把分数变成概率。比如区分猫、狗、卡车,一张“卡车”的图片,网络可能输出 [0.7, 0.2, 0.1],意思是 70% 像猫,20% 像狗,10% 像卡车(这里假设类别顺序是猫、狗、卡车)。

明白 这个分布里藏着什么。真实的标签是“猫”,但网络同时给出了“像狗”和“像卡车”的概率。这 20% 和 10% 就是“暗知识”。它告诉学生:在教师看来,猫和狗比猫和卡车更相似。这种相似性关系,是原始标签(one-hot,如 [1, 0, 0])完全无法提供的。

对比 原始标签与教师分布。直接拿 one-hot 标签训练学生,学生只会学到“猫就是猫,和其他类别无关”。而拿教师的分布训练,学生能学到“猫和狗有相似特征,和卡车差别较大”。这就是蒸馏能提升小模型性能的根本原因。


2. 引入温度参数软化分布

理解 原始 Softmax 的问题。标准的 Softmax 公式是:给定第 $i$ 类的分数 $z_i$,其概率 $p_i$ 为:

$$p_i = \frac{e^{z_i}}{\sum_{j} e^{z_j}}$$

这个公式会把分数差异放大。如果教师网络对“猫”的分数是 10,对“狗”的分数是 9,那么 Softmax 输出的概率大约是 [0.73, 0.27],而这个分布已经相当尖锐(靠近 one-hot),暗知识被压缩得很小,学生很难从中学习。

引入 温度系数 $T$。在 Softmax 计算中,将每个分数除以 $T$,公式变为:

$$p_i = \frac{e^{z_i / T}}{\sum_{j} e^{z_j / T}}$$

观察 温度 $T$ 的效果。当 $T = 1$ 时,就是标准 Softmax。当 $T > 1$ 时(例如 $T = 4$),分数差异被缩小,输出的概率分布变得更平缓,小概率类别的相对差异被放大,这就是“软化”。当 $T \to \infty$ 时,分布趋近于均匀分布。当 $T < 1$ 时,分布变得更尖锐,$T \to 0$ 时趋近于 one-hot

记住 硬标签与软标签的命名。教师网络在高温(如 $T = 4$)下输出的分布,叫软标签soft label)。真实标签(one-hot)叫硬标签hard label)。


3. 计算软标签之间的KL散度

设定 场景。教师网络和学生网络,都使用相同的温度 $T$ 计算软标签。设教师网络的软标签分布为 $P$,学生网络的软标签分布为 $Q$。

使用 KL散度公式。KL散度衡量两个分布之间的差异,公式为:

$$KL(P \| Q) = \sum_{i} P(i) \log \frac{P(i)}{Q(i)}$$

解读 这个公式的直觉。如果 $P(i) = Q(i)$,则 $\log \frac{P(i)}{Q(i)} = \log 1 = 0$,这一项贡献为 0。如果 $P(i) > Q(i)$,则对数大于 0,产生正惩罚;如果 $P(i) < Q(i)$,对数为负,但因为前面乘以 $P(i)$,惩罚可能为负(即模型“过于自信”地低估了某个教师认为可能的小概率类别,会得到负梯度,推动学生提高该类别概率)。

注意 KL散度的不对称性。$KL(P \| Q)$ 通常不等于 $KL(Q \| P)$。在蒸馏中,通常固定教师分布 $P$,优化学生分布 $Q$,让 $Q$ 去贴合 $P$。

对比 交叉熵与KL散度。训练学生网络时,如果直接最小化 $KL(P \| Q)$,等价于最小化 $P$ 和 $Q$ 的交叉熵(因为 $KL(P \| Q) = H(P, Q) - H(P)$,而 $H(P)$ 是常数,优化中可忽略)。


4. 组合硬标签损失与软标签损失

构建 总损失函数。实践中,学生网络的训练目标通常包含两项:

第一项:软标签损失。计算教师软标签 $P_T$ 与学生软标签 $Q_T$ 的KL散度,乘以温度平方 $T^2$。公式为:

$$L_{soft} = T^2 \cdot KL(P_T \| Q_T)$$

为什么 要乘以 $T^2$?因为 $Q_T$ 的梯度大小与 $1/T^2$ 成比例(对 $z_i$ 求导后约简)。为了保持不同温度下梯度幅度一致,乘回 $T^2$。

第二项:硬标签损失。计算学生网络在 $T = 1$ 时的输出与真实标签 one-hot 的交叉熵:

$$L_{hard} = CE(y, \sigma(z_s))$$

其中 $\sigma$ 是标准 Softmax,$z_s$ 是学生网络的 logits。

组合 两项损失:

$$L = \alpha \cdot L_{hard} + (1 - \alpha) \cdot L_{soft}$$

设置 超参数 $\alpha$。通常取 $\alpha = 0.7$ 左右。如果 $\alpha$ 过大,学生过度关注硬标签,忽略暗知识;如果 $\alpha$ 过小,学生可能无法学对“最终答案”,准确率下降。

选择 温度 $T$。常用取值在 2 到 8 之间。更大的 $T$ 让分布更平滑,暴露更多暗知识,但也可能引入过多噪声。需要根据任务调节。


5. 实际训练步骤

准备 预训练的教师网络。确保教师网络在目标任务上有良好性能。

固定 教师网络参数。训练过程中,教师网络不更新,只用于前向传播,产生软标签。

加载 学生网络和优化器。

执行 循环训练:

  1. 一个小批次数据 (x, y)
  2. 计算 教师软标签:用温度 $T$ 计算教师网络的 Softmax 输出 $P_T$。
  3. 计算 学生软标签:用相同温度 $T$ 计算学生网络的 Softmax 输出 $Q_T$。
  4. 计算 学生硬标签输出:用 $T = 1$ 计算学生网络的 Softmax 输出 $\hat{y}$。
  5. 计算 损失 $L = \alpha \cdot CE(y, \hat{y}) + (1 - \alpha) \cdot T^2 \cdot KL(P_T \| Q_T)$。
  6. 反向传播 更新学生网络参数。

重复 步骤 1 到 6 直到收敛。

调整 温度 $T$。可以在验证集上尝试不同温度(如 2468),选择效果最好的值。


6. 一个核心公式总结

整合 所有元素,知识蒸馏的完整损失函数为:

$$L_{KD} = \alpha \cdot CE(y, \sigma(z_s)) + (1 - \alpha) \cdot T^2 \cdot KL\left(\sigma\left(\frac{z_t}{T}\right) \| \sigma\left(\frac{z_s}{T}\right)\right)$$

其中 $z_t$ 是教师网络的 logits,$z_s$ 是学生网络的 logits。

理解 这个公式的两层含义。第一层:学生要模仿教师的最终决定(硬标签)。第二层:学生要模仿教师的思考过程(软标签分布)。温度 $T$ 控制“思考过程”暴露的细节程度,KL散度量化“思考过程”的差距。

把握 关键直觉。温度太高,分布太平,学生学到的都是“什么都像一点”的模糊知识;温度太低,分布太尖,学生只知道“最像哪个”,丢失了暗知识。合适的中等温度(如 4)既能保留相似性关系,又能提供足够的梯度信号。

最终 知识蒸馏的KL散度与温度软化,本质上是通过调节分布形态,把教师网络中“如何看世界”的信息编码成梯度,逐步传递给学生网络。

评论 (0)

暂无评论,快来抢沙发吧!

扫一扫,手机查看

扫描上方二维码,在手机上查看本文