从 RL for LLM 视角重新理解 KL 估计:读《Approximating KL Divergence》笔记

发布于
8分钟
3,106
en / zh语言

PPO 和 GRPO 中使用的 KL 散度估计算法有什么不同?为什么会这样?

John Schulman 的博客《Approximating KL Divergence》讲述了如何用采样(Monte Carlo)近似 KL 散度,并介绍了三种估计器(k1k_1k2k_2k3k_3)及它们的偏差与方差特性。但原文是在通用概率分布背景下讨论的,没有涉及 LLM 的强化学习训练场景。这篇文章记录了我在阅读原文时产生的疑问、结合 RL for LLM 的思考,以及对原文某些解释的质疑。

《Approximating KL Divergence》 讲了什么

在这一个章节,我们带着没有读过这篇文章的读者,快速过一遍这篇文章中讲的最重要的内容。简单得说,这篇文章讨论了在无法直接计算出 KL 散度时,我们如何基于蒙特卡洛采样的方式构造合理的 KL 散度的优化器。

KL(q,p)=xq(x)logq(x)p(x)=Exq[logq(x)p(x)]\text{KL}(q, p) = \sum_x q(x) \text{log}\frac{q(x)}{p(x)} = \mathbb{E}_{x\sim q}\left [\text{log}\frac{q(x)}{p(x)}\right ](1)

如上式所示:在估计两个(复杂)分布的 KL 散度的时候,代码中常用的一个技巧是:将 KL 散度近似为 log(q(x)p(x))\text{log}(\frac{q(x)}{p(x)}) 的样本均值(样本来自 qq,而不是更标准的 Eq[log(q(x)x(x))]\mathbb{E}_q [\text{log}(\frac{q(x)}{x(x)})]。用 12(logr)2\frac{1}{2} (\log r)^2 的样本平均代替更标准的 logr\log r 形式。本文将解释为什么这个表达式是一个不错(尽管有偏)的 KL 估计器,以及如何在保持其低方差的同时使它无偏。

我们能计算 KL 的方法取决于我们对 ppqq 的访问方式。这里我们假设可以计算任意 xx 的概率(或概率密度)p(x)p(x)q(x)q(x),但不能解析地对 xx 求和(或积分)。为什么会不能解析求和呢?有可能精确计算需要过多的计算量或内存,没有解析解,或者我们只存储 log-prob 而不是整个分布,这能简化代码。这是合理的选择,尤其当 KL 只是作为诊断使用时(例如在强化学习中常见的情况)。最常见的求和或积分估计策略是 蒙特卡罗估计。给定样本 x1,x2,,xnqx_1, x_2, \dots, x_n \sim q,我们如何构造一个好的估计器?

一个好的估计器应当无偏(均值正确)且低方差。我们知道一个无偏估计器(PPO 的选择)是:

k1=logq(x)p(x)k_1 = \log \frac{q(x)}{p(x)}(2)

但它的方差很高,因为KL散度就定义而言一定是一个正值,但是对于上边的估计来说,它的一半的样本是负的(假设我们在没有任何对 ppqq 的先验知识下),所以他的方差很高。为了简化表达,我们记:r=q(x)p(x)r = \frac{q(x)}{p(x)},那么原始的 KL 散度的估计可以写为:

KL[q,p]=Exq[logr]\mathrm{KL}[q, p] = \mathbb{E}_{x \sim q} [\log r](3)

为了减少方差,我们可以设计另外一个替代的 estimator:

k2=12(logr)2k_2 = \frac{1}{2} (\log r)^2(4)

它的方差低,但有偏。直觉上来说,k2k_2 更好,因为每个样本都给出了 ppqq 之间的“距离”,因而它总是正的。经验上来说,k2k_2 的方差确实远低于 k1k_1,而且偏差也小了很多很多。至于为什么 k2k_2k1k_1 的方差要低很多,原文中,利用f-散度给出了一个分析性的解释,在这里不再赘述。

那么有没有可能构造一个无偏并且方差又小的 KL 散度的估计器呢?一个降低方差的一般方法是使用控制变量(control variate):我们已知 k1k_1 是一个无偏估计,那么我们取 k1k_1,加上一个期望为零、且与其负相关的项,这样就能减小方差。我们能够轻而易举得到的一个期望为0的量是 r1r-1。于是,对于任意 λ\lambda

k=logr+λ(r1)k = -\log r + \lambda (r - 1)(5)

都是无偏的 KL 估计器。虽然理论上我们可以通过优化来最小化方差并解出最优 λ\lambda,但结果依赖于 p,qp, q,难以解析求出。注意到由于 log(x)\log(x) 是凹的:

log(x)x1\log(x) \le x - 1(6)

因此若取 λ=1\lambda = 1,上式总为正。这里,r1r - 1 其实是 logr\text{log}rr=1r = 1 处的切线。所以当 λ=1\lambda=1, 我们其实是在测量 log(x)\log(x) 与其切线的垂直距离。所以我们可以得到估计器:

k3=(r1)logrk_3 = (r - 1) - \log r(7)

所以它总是正的。而 k3k_3 正是 GRPO 在 KL 散度估计上区别于 PPO 的一大特点 (PPO 使用的是 k1k_1)。

在 RL for LLM 的角度讨论 KL 散度的估计

在 RL(比如 PPO、GRPO 等)里,我们经常会在 loss 里加一个 KL 散度,用来约束新策略(policy)不要偏离旧策略太多。这里,qq 是旧策略分布(πold\pi_{\text{old}},而 pp 是新策略分布(πnew\pi_{\text{new}}),xx 则是一个完整的动作样本(在 LLM 里就是 token 或 token 序列)。我们通常用 ss 表示 state(在 LLM 里就是 prompt 或上下文), xx 就是在这个上下文下生成的一个具体 token。在计算 KL 时,我们实际上是对 给定状态下的动作分布 做 KL,然后再对状态取期望:

KL[p,q]=Es[xp(xs)logp(xs)q(xs)]\mathrm{KL}[p, q] = \mathbb{E}_{s} \left[ \sum_x p(x|s) \log \frac{p(x|s)}{q(x|s)} \right](8)

在采样时,我们通常固定一个 prompt(state),然后针对它去估计这个 KL 值。

那么为什么无法直接计算 KL 的值,而一定要估计呢? 原因正如原文提到的三条,在 RL for LLM 里主要是 第 1 条:“动作空间(token space)太大,无法对所有可能的 xx 枚举求和/积分。” 比如一个 tokenizer 有 50,000 个词表,即便只计算一个 token 的 KL,都要对这 50,000 个动作求和;而 RL 中通常是多步生成(sequence),这个空间就会是指数级的,因而根本不可行。另外,还有一个现实原因:在训练时,我们通常没有直接保存完整的分布(所有 token 的概率),只保存了生成路径上选中的 token 的 log-prob。这是为了节省显存和 I/O 开销。因此我们只能用 蒙特卡洛采样:从某个分布(通常是 qq,即旧策略)采样一些 xx,用这些样本来近似 KL。这里就回归到了上边这个博客所讨论的范畴。

在上边这篇文章里,我们一直在讨论的估计器(estimator),其实就是一个基于采样的函数,它输入的是在某个样本 xxp(x)p(x)q(x)q(x) 的值(比如他们的比值 r=q(x)p(x)r = \frac{q(x)}{p(x)}),输出一个数,我们用这个数的样本均值用于近似 KL 散度。比如:

  • k1(x)=logrk_1(x) = -\log r
  • k2(x)=12(logr)2k_2(x) = \frac12 (\log r)^2
  • k3(x)=(r1)logrk_3(x) = (r - 1) - \log r

这些 kik_i 就是不同的形式的 KL-estimator,它们都能通过在采样上取平均来近似 KL 散度,只是无偏性、方差大小不同。当我们选定了估计器,我们其实就决定用哪种公式去近似 KL。操作过程是:

  1. 采样 从旧策略 qq 采样一批 token(或序列)x1,x2,,xNx_1, x_2, \dots, x_N
  2. 计算 log-prob 对每个样本,计算新策略和旧策略下的对数概率:
logp(xi), logq(xi)\log p(x_i),\ \log q(x_i)(9)

并求出 ri=q(xi)p(xi)r_i = \frac{q(x_i)}{p(x_i)}logri\log r_i

  1. 代入 estimator 公式 比如选 k3k_3,就算:
k3(xi)=(ri1)logrik_3(x_i) = (r_i - 1) - \log r_i(10)
  1. 求平均
KL^1Ni=1Nk3(xi)\widehat{\mathrm{KL}} \approx \frac1N \sum_{i=1}^N k_3(x_i )(11)

这就是 KL 的近似值,用它替代真实 KL。

如果我们对比 KL 散度的真实计算过程(不估计),对于离散型概率分布(LLM 单步 token):我们要遍历每一个可能的token xx

KL(pq)=xp(x)logp(x)q(x)\mathrm{KL}(p\|q) = \sum_x p(x) \log \frac{p(x)}{q(x)}(12)

这里可以通过对比看出,使用 estimator 后计算量其实是巨大的。

关于不同 KL Estimator 的方差的讨论

需要指出的是,这里我们所说的方差,是说通过估计器采样出来的值的方差

Varxq[k(x)]\mathrm{Var}_{x \sim q}[k(x)](13)

即估计器 k(x)k(x) 在样本空间上的波动程度。一个无偏的估计器,是指在无限多样本下,它的均值等于真实 KL;那么一个方差大的估计器,就可能造成即便期望正确(无偏),少量样本的平均值也可能偏离很远。在 RL for LLM 中,KL 项通常是 loss 里的一个正则化约束系数(比如 βKL\beta \cdot \mathrm{KL}),如果 KL 估计器的方差大,就会导致 loss 抖动,进而导致梯度抖动、训练不稳定。

原文中,作为为了让读者直观地理解为什么 k1k_1 不是一个小方差的估计器,给出了如下解释:

However, it (k1k_1) has high-variance, as it’s negative for half of the samples, whereas KL is always positive.

作者指出,k1k_1 虽然是无偏估计,但是由于 ppqq 在无先验知识的条件下,采样出来的样本值总是一半比对方大,一半比对方小,所以导致 k1k_1 采集到的数据总是一部分大于 00 另一半小于 00,到这里我认为都是没问题的,但是作者后边说:因为 KL 总是大于 0 的 (可由基本不等式推出),所以 k1k_1 的方差很高。这里因果关系不成立的原因在于,我们不能用期望值的正负去要求估计值采样的正负。一个简单的反驳是,如果计算期望,p(x)logp(x)q(x)p(x) \log \frac{p(x)}{q(x)} 也是一半是正一半是负,这并不能说明什么。其实单次采样的 log 比值(无论是 logq(x)p(x)\log \frac{q(x)}{p(x)} 还是 logp(x)q(x)\log \frac{p(x)}{q(x)})确实都可能正也可能负,和 k1k_1 一样会“跳变”,所以正负跳变本身并不是导致方差大的唯一原因

根据 KL 散度的定义:

KL(qp)=Exq[logq(x)p(x)]\mathrm{KL}(q \| p) = \mathbb{E}_{x\sim q}\left[ \log \frac{q(x)}{p(x)} \right](14)

这个期望是必然非负,但里面的被积项 logq(x)p(x)\log\frac{q(x)}{p(x)} 在单个样本上可以正也可以负。k1k_1 其实就是直接取这个被积项:

k1(x)=logq(x)p(x)k_1(x) = \log \frac{q(x)}{p(x)}(15)

所以它的每个样本值确实可能正也可能负,和 KL 定义中的 integrand 一样。

那为什么 k1k_1 方差大?

不是因为“正负跳变”本身,而是因为 k1k_1 的取值分布通常很宽(heavy-tailed),例如,在概率很小的样本 p(x)p(x) 下,logqp\log\frac{q}{p} 可以变得非常大(正或负)。这种极端值在有限样本平均中占据很大权重,所以会导致高方差。换句话说,问题是极端值 + 正负抵消的组合:正负抵消让你需要更多样本才能收敛到真实均值,极端值让样本方差本身更大。所以文章里提的“因为一半是负的”只是一个直观类比,并不是完整的方差来源解释。

如果我们从这个角度,来看其他估计器 k2k_2k3k_3,就会发现:k2=12(logr)2k_2 = \frac12 (\log r)^2 永远为正,去掉了正负抵消的可能,因为带来了偏差;但是引入了平方,使得数值变得更平滑,使方差下降;k3k_3 通过控制变量消掉了一部分波动源,在保留无偏性的同时降低了方差,具体请看下边的解释。

在 PPO/GRPO 里,如果用 k1k_1,当 batch 小、分布差距大时,KL 估计值会很抖动(因为几个极端样本能显著改变均值),同时会使 KL-penalty 系数调节不稳(可能突然过强或过弱),换成低方差估计器(k2k_2k3k_3)时,batch 中的每个样本对 KL 的贡献更稳定,不容易被少数极端样本主导。

为什么 k3k_3 能做到无偏且小方差?

乍一看,k3k_3 全是正数,直觉会以为它的均值要大于 k1k_1 的均值。
但别忘了 k3k_3 是通过 控制变量(control variate)k1k_1 推出来的。文章中给的思路是:

k~(x)=k1(x)+λh(x)\tilde{k}(x) = k_1(x) + \lambda \cdot h(x)(16)

其中 h(x)=r1h(x) = r - 1,在 xqx\sim q 下的期望是:

Exq[h(x)]=Eq[p(x)q(x)1]=xp(x)1=11=0\mathbb{E}_{x\sim q}[h(x)] = \mathbb{E}_q\left[\frac{p(x)}{q(x)} - 1\right] = \sum_x p(x) - 1 = 1 - 1 = 0(17)

所以加上任何倍数的 h(x)h(x),不会改变期望值。当取 λ=1\lambda = 1 时:

k~(x)=logr+(r1)=(r1)logr=k3(x)\tilde{k}(x) = -\log r + (r - 1) = (r - 1) - \log r = k_3(x)(18)

这就解释了为什么 k3k_3 的期望等于 k1k_1 的期望,也等于 KL,即是一个无偏估计。

k3k_3k1k_1 的方差小的原因在于:k1k_1 里只有 logr-\log r,波动可能很大(正负都有,大值出现时特别大);而 r1r - 1logr-\log r 在数值上高度相关(一个变大时另一个也变大/变小),而且相关性是负的,加上 (r1)(r - 1) 相当于用一个负相关的量去抵消波动。抵消之后剩下的 k3k_3,取值范围收紧,且全为正值,最终会导致样本方差下降。

引用
@misc{zhang2025kl,
  author = {Zhang, Yihua},
  title  = {从 RL for LLM 视角重新理解 KL 估计:读《Approximating KL Divergence》笔记},
  year   = {2025},
  url    = {https://normaluhr.github.io/zh/p/kl}
}