文章

正向/反向 KL散度 与 模型蒸馏

正向/反向 KL散度 与 模型蒸馏

问题:假设目标分布 $p$ 已知,我们希望用一个受限的分布族 $q$ 去逼近它,那么究竟应该最小化 $D_{\mathrm{KL}}(p|q)$,还是 $D_{\mathrm{KL}}(q|p)$?

t2

题图给出了这个问题最经典的可视化例子:$p$ 是一个双峰高斯混合,而 $q$ 被限制成单个高斯。使用正向 KL 时,最优的 $q$ 倾向于把两个峰一起“包住”;使用反向 KL 时,最优的 $q$ 则倾向于选择其中一个峰。

这一差异在大语言模型的知识蒸馏中尤其重要。MiniLLM 明确指出,传统知识蒸馏通常近似最小化 teacher 到 student 的正向 KL,它要求 student 尽量覆盖 teacher 的所有模式;但开放式文本生成的输出空间高度多模态,而容量较小的 student 往往无法忠实表示 teacher 的全部模式,因此这种覆盖行为可能使 student 在 teacher 的低概率区域赋予过多概率质量。MiniLLM 因而转向反向 KL,希望 student 聚焦于 teacher 的主要模式并减少对低概率区域的覆盖。

信息论

t1

信息量与最优编码长度

假设一个离散随机变量真实服从分布 $p(x)$。当事件 $x$ 发生时,其信息量可以写成 $-\log p(x)$。若对数以 $2$ 为底,这个量可以直接解释为:在理想编码条件下,为事件 $x$ 分配的编码长度大约是多少 bit。

因此,如果数据真的来自 $p$,使用针对 $p$ 优化的编码方案时,平均编码长度就是熵

\[H(p)=-\mathbb{E}_{x\sim p}\log p(x) =-\sum_x p(x)\log p(x).\]

这可以理解为分布 $p$ 本身所固有的平均信息量,也对应在理想条件下对来自 $p$ 的数据进行无损压缩时能够达到的平均编码长度下界。

但现在假设真实数据仍来自 $p$,我们并不知道 $p$,而是按照另一个分布 $q$ 来设计编码器。事件 $x$ 会被分配大约 $-\log q(x)$ 的编码长度,于是平均长度变成

\[H(p,q) = -\mathbb{E}_{x\sim p}\log q(x).\]

这就是交叉熵。现实按照 $p$ 产生数据,但编码器按照 $q$ 分配码长时,平均要付出多少编码成本?

因此,固定 $p$ 而优化 $q$ 时,交叉熵在 $q=p$ 处达到最小值。这也是交叉熵能够自然成为最大似然学习目标的原因。

KL 散度

将交叉熵减去理论上最优的编码长度 $H(p)$,便得到

\[D_{\mathrm{KL}}(p\|q) = H(p,q)-H(p).\]

展开以后就是熟悉的正向 KL:

\[D_{\mathrm{KL}}(p\|q) = \mathbb{E}_{x\sim p} \left[ \log\frac{p(x)}{q(x)} \right].\]

因此,在信息论视角下,$D_{\mathrm{KL}}(p|q)$ 有一个非常清晰的含义:真实数据来自 $p$,但我们错误地使用 $q$ 进行编码时,相比使用最优编码器平均额外浪费了多少信息。

因为在优化 $q$ 时 $H(p)$ 与参数无关,所以

\[\arg\min_q D_{\mathrm{KL}}(p\|q) = \arg\min_q H(p,q).\]

机器学习里常见的 cross-entropy loss,本质上正是在最小化一个正向 KL,只不过与模型参数无关的 $H(p)$ 被省略了。

正向KL和反向KL

现在分别写出两个方向:

\[D_{\mathrm{KL}}(p\|q) = \mathbb{E}_{x\sim p} \left[ \log\frac{p(x)}{q(x)} \right],\]

以及

\[D_{\mathrm{KL}}(q\|p) = \mathbb{E}_{x\sim q} \left[ \log\frac{q(x)}{p(x)} \right].\]

表面上看似乎只是把 $p$ 和 $q$ 交换了一下,但真正发生变化的是期望所使用的采样分布。

正是这一点造成了两种完全不同的归纳偏置。

正向 KL:模式覆盖

考虑 $D_{\mathrm{KL}}(p|q)$。由于 $x\sim p$,只要某一区域在 $p$ 下拥有明显概率质量,它就会频繁出现在目标函数的期望中。

如果某处满足 $p(x)>0$,但模型给出 $q(x)\rightarrow0$,那么

\[\log\frac{p(x)}{q(x)} \rightarrow +\infty.\]

因此,正向 KL 对“真实分布认为可能,但模型认为几乎不可能”的情况非常敏感。换句话说,$q$ 很难直接放弃 $p$ 的某个重要模式。

这就是所谓的 mode-covering 行为:与其完全漏掉一个峰,优化器往往宁愿把 $q$ 拉宽一些,将多个峰全部覆盖进去。

反向 KL:模式聚集

反过来看 $D_{\mathrm{KL}}(q|p)$。此时期望中的样本来自 $q$。因此,只要 $q$ 经常生成某一区域的样本,这些位置就会进入损失。

如果 $q(x)>0$,而目标分布满足 $p(x)\rightarrow0$,那么

\[\log\frac{q(x)}{p(x)} \rightarrow+\infty.\]

因此,反向 KL 强烈反对 student 在 teacher 认为不合理的区域放置概率质量。相比之下,如果 $p$ 在另外某个峰上还有大量概率,但 $q$ 根本不到那里,那么这些区域因为很少被 $q$ 采样,反而不会对期望产生直接而强烈的惩罚。

所以反向 KL 的典型策略变成mode-seeking:与其勉强覆盖所有模式并穿过中间的低概率区域,不如选择一个高密度模式,把概率集中在那里。

双高斯混合为什么产生完全不同的最优解

解析计算

考虑最简单的对称双峰分布

\[p(x) = \frac{1}{2}\mathcal N(x;-a,\sigma^2) + \frac{1}{2}\mathcal N(x;a,\sigma^2),\]

而近似分布被严格限制为单个高斯

\[q(x)=\mathcal N(x;\mu,s^2).\]

问题:在 $q$ 无法完整表达双峰结构时,不同 KL 会让它以什么方式妥协。

t2

最小化正向 KL

首先考虑

\[q_F^\star = \arg\min_qD_{\mathrm{KL}}(p\|q).\]

因为 $H(p)$ 与 $q$ 无关,所以只需要最小化 $-\mathbb E_p\log q(x)$。对于高斯 $q$,

\[-\log q(x) = \frac{1}{2}\log(2\pi s^2) + \frac{(x-\mu)^2}{2s^2}.\]

代入期望后,有

\[D_{\mathrm{KL}}(p\|q) = C + \frac{1}{2}\log s^2 + \frac{1}{2s^2} \mathbb E_{x\sim p}(x-\mu)^2,\]

其中 $C$ 与 $\mu,s$ 无关。于是最优参数满足

\[\mu^\star=\mathbb E_p[x], \qquad (s^\star)^2=\operatorname{Var}_p(x).\]

对于上面的对称混合,

\[\mathbb E_p[x]=0, \qquad \operatorname{Var}_p(x)=a^2+\sigma^2.\]

因此

\[q_F^\star = \mathcal N(0,a^2+\sigma^2).\]

结果:正向 KL 并不会选择左峰或者右峰,而是产生一个位于两个峰之间、方差很大的高斯。

追求:两个峰都是 $p$ 真正会产生数据的位置,任何一个都不能漏掉。既然单个高斯没有能力同时形成两个峰,那就只能把方差增大,将两个区域都覆盖。

代价:这个高斯不可避免地会在两个峰之间的“山谷”区域赋予相当大的概率,而原始混合分布在那里可能几乎没有概率质量。

MiniLLM 所讨论的 LLM 蒸馏问题与这一 toy example 完全对应:当 student 的表示能力不足时,正向 KL 会迫使它试图覆盖复杂 teacher 分布中的多个模式,并可能因此在 teacher 的低概率区域产生不合理的概率质量。

最小化反向 KL

现在考虑

\[q_R^\star = \arg\min_qD_{\mathrm{KL}}(q\|p).\]

将式子展开:

\[D_{\mathrm{KL}}(q\|p) = \mathbb E_q[\log q(x)] - \mathbb E_q[\log p(x)].\]

或者写成

\[D_{\mathrm{KL}}(q\|p) = -H(q)-\mathbb E_q[\log p(x)].\]

如果让单个高斯同时覆盖左右两个峰,那么它不可避免地会从两个峰之间采样大量样本。但在这些样本上 $p(x)$ 很小,于是 $-\log p(x)$ 很大,从而受到明显惩罚。

相反,如果 $q$ 只落在右侧模式附近,那么它生成的绝大部分样本在 $p$ 看来都是合理的。对于充分分离的两个高斯,即 $a\gg\sigma$ 时,在右峰附近可以近似写成

\[p(x) \approx \frac12\mathcal N(x;a,\sigma^2).\]

如果令

\[q(x)\approx\mathcal N(x;a,\sigma^2),\]

那么局部上近似有

\[\frac{q(x)}{p(x)} \approx 2,\]

因而

\[D_{\mathrm{KL}}(q\|p) \approx\log 2.\]

左侧模式完全类似。

因此在对称问题中通常存在两个等价解:一个选择左峰,一个选择右峰。具体优化最终落在哪一个峰上,可能由初始化、采样噪声或者数值误差打破对称性。

知识蒸馏

Forward KL 蒸馏

假设 teacher 分布为 $p(y\mid x)$,student 为 $q_\theta(y\mid x)$。标准蒸馏可写成

\[D_{\mathrm{KL}} \left( p(\cdot\mid x) \| q_\theta(\cdot\mid x) \right).\]

展开后,

\[D_{\mathrm{KL}}(p\|q_\theta) = \mathbb E_{y\sim p} \left[ \log p(y\mid x) - \log q_\theta(y\mid x) \right].\]

teacher 固定以后,第一项与 student 参数无关,因此优化本质上仍然是

\[-\mathbb E_{y\sim p} \log q_\theta(y\mid x).\]

换言之,这是用 teacher 的概率质量来加权 student 的 cross-entropy。teacher 认为有可能的输出,student 都需要尽量给出足够概率。

对于分类问题,这通常没有明显矛盾,因为类别数有限,student 完全可能覆盖 teacher 的主要概率质量。但 MiniLLM 指出,在开放式语言生成中,teacher 的条件分布可能包含大量语义、措辞、长度和推理路径不同的模式,而小 student 的容量有限;要求它覆盖所有模式可能并不是合适的训练偏置。

Reverse KL蒸馏:student 自己生成,然后接受 teacher 检验

t3

MiniLLM 改为最小化

\[\mathcal L(\theta) = D_{\mathrm{KL}} (q_\theta\|p).\]

更完整地,

\[\mathcal L(\theta) = \mathbb E_{x} \mathbb E_{y\sim q_\theta(\cdot\mid x)} \left[ \log \frac{q_\theta(y\mid x)} {p(y\mid x)} \right].\]

$y$ 现在由 student 自己生成。训练所关注的区域因此变成 student 真正会在推理阶段访问的区域。

MiniLLM 将这种目标称为 on-policy distillation,并通过 policy-gradient 形式求解;论文同时指出,为控制策略梯度中的方差、reward hacking 和长度偏差,又引入了 single-step decomposition、teacher-mixed sampling 和 length normalization。

Reverse KL 等价形式:teacher reward 加 student entropy

反向 KL 展开后有

\[D_{\mathrm{KL}}(q_\theta\|p) = \mathbb E_{q_\theta}\log q_\theta(x) - \mathbb E_{q_\theta}\log p(x).\]

利用 $H(q_\theta)=-\mathbb E_{q_\theta}\log q_\theta(x)$,可以写成

\[D_{\mathrm{KL}}(q_\theta\|p) = -H(q_\theta) - \mathbb E_{q_\theta}\log p(x).\]

所以最小化 reverse KL 等价于

\[\max_\theta \left[ \mathbb E_{x\sim q_\theta}\log p(x) + H(q_\theta) \right].\]
  • 第一项要求 student 生成的结果在 teacher 下拥有较高概率;
  • 第二项则奖励 student 保留自身熵,避免简单塌缩成一个确定性输出。

reverse KL 本身已经具有一种 entropy-regularized reward maximization 的结构。

这正是 MiniPLM 给出的另一个解释。其论文将 reverse KL 改写为 reward maximization,并定义

\[r(p,q_\theta,x) = \log \frac{p(x)}{q_\theta(x)}.\]

于是

\[\min_\theta D_{\mathrm{KL}}(q_\theta\|p) = \max_\theta \mathbb E_{x\sim q_\theta} r(p,q_\theta,x).\]

论文对此的解释是:较大的 $\log p(x)$ 表示 teacher 偏好该文本,而 $q_\theta$ 项同时与输出多样性相关。

MiniPLM:先用 reverse-KL 思想做数据增强,再标准训练

t4

MiniPLM 对这个问题做了进一步转换。直接进行 reverse-KL on-policy 蒸馏需要 student 在线采样,还需要 teacher 在线计算概率,这对于大规模预训练成本很高。MiniPLM 因而使用一个较小的 reference model $p_{\mathrm{ref}}$,把样本的“价值”近似写成

\[r(p,p_{\mathrm{ref}},x) = \log \frac{p(x)} {p_{\mathrm{ref}}(x)}.\]

直观上,如果 teacher 对某个样本给出高概率,但较弱的 reference model 给出低概率,那么该样本很可能包含“小模型尚未掌握、而大模型已经掌握”的困难信息。Difference Sampling 因此利用这一概率比值从原始语料中筛选和重加权训练数据。

有意思的是,在完成这种基于概率差异的数据筛选以后,student 并不是继续执行复杂的 KL 或 policy-gradient 训练,而是重新回到标准 next-token cross-entropy:

\[\mathcal L(q_\theta,D') = - \frac{1}{|D'|} \sum_{x\in D'} \frac{1}{|x|} \sum_{t=1}^{|x|} \log q_\theta(x_t\mid x_{<t}).\]

MiniPLM_method

也就是说,MiniPLM 可以概括成一种两阶段思想:先使用 reverse-KL / reward 的思想改变“什么数据值得训练”,然后使用标准 cross-entropy 完成“怎样学习这些数据”。 论文明确描述了从 Difference Sampling 构造 $D’$,再在 $D’$ 上以 cross-entropy 从头预训练 student 的流程。

DPKD:只选 KL 的方向还不够

t5

论文 DPKD 进一步指出,单纯的 KL divergence 可能仍不足以描述 LLM 蒸馏所关心的输出质量和偏好结构。论文首先定义了常见的 forward-KL 和 reverse-KL,然后在 reverse KL 之外引入 implicit reward,并进一步把 teacher output 与 student output 组织成 preference optimization 问题。

其基础目标可以抽象写为

\[\max_\theta \mathbb E \left[ r_p(y\mid x) - \beta D_{\mathrm{KL}} \bigl( q_\theta(y\mid x)\|p(y\mid x) \bigr) \right].\]

这里 $D_{\mathrm{KL}}(q_\theta|p)$ 起到 regularization 的作用,而 $r_p$ 则承担对输出质量进行额外刻画的角色。随后借助 Bradley-Terry preference formulation,将 teacher 输出 $y_t$ 和 student 输出 $y_s$ 的相对偏好转化成一个 logistic objective,并进一步加入长度归一化和语言模型损失。

总结

对于一般问题,不应把选择原则表述成“哪个 KL 更好”,而应该先明确最终损失究竟希望控制什么。

如果真正关心的是:来自 $p$ 的样本不能被 $q$ 漏掉,那么 $D_{\mathrm{KL}}(p|q)$ 是自然的目标。标准最大似然、cross-entropy 训练以及很多传统 teacher-to-student distillation 都属于这一逻辑。它会强烈惩罚 $q$ 对真实高概率区域给出过小概率,因此倾向于 coverage。当 $q$ 有足够容量表达 $p$ 时,这通常也是非常自然的选择。

如果真正关心的是:从 $q$ 生成出来的样本不能跑到 $p$ 的低概率区域,那么 $D_{\mathrm{KL}}(q|p)$ 更直接。它从 $q$ 自己的分布中采样并使用 $p$ 检验这些样本,因此尤其适合描述“生成出来的内容是否得到 teacher 认可”这一问题。当 $q$ 的容量远小于复杂的多模态 $p$ 时,它通常会牺牲一部分 mode coverage,换取更集中的高概率生成区域。

Forward KL 在意的是 $p$ 产生什么;Reverse KL 在意的是 $q$ 会产生什么。

生成学习中一个非常基础的设计选择:当模型容量不足以完整复制目标分布时,我们究竟希望它覆盖所有可能性,还是希望它只生成自己最有把握、同时真实分布也高度认可的可能性?

warning

本文由作者按照 CC BY 4.0 进行授权