从 RL for LLM 视角重新理解 KL 估计:读《Approximating KL Divergence》笔记
PPO 和 GRPO 中使用的 KL 散度估计算法有什么不同?为什么会这样?
John Schulman 的博客《Approximating KL Divergence》讲述了如何用采样(Monte Carlo)近似 KL 散度,并介绍了三种估计器(、、)及它们的偏差与方差特性。但原文是在通用概率分布背景下讨论的,没有涉及 LLM 的强化学习训练场景。这篇文章记录了我在阅读原文时产生的疑问、结合 RL for LLM 的思考,以及对原文某些解释的质疑。
《Approximating KL Divergence》 讲了什么
在这一个章节,我们带着没有读过这篇文章的读者,快速过一遍这篇文章中讲的最重要的内容。简单得说,这篇文章讨论了在无法直接计算出 KL 散度时,我们如何基于蒙特卡洛采样的方式构造合理的 KL 散度的优化器。
如上式所示:在估计两个(复杂)分布的 KL 散度的时候,代码中常用的一个技巧是:将 KL 散度近似为 的样本均值(样本来自 ,而不是更标准的 。用 的样本平均代替更标准的 形式。本文将解释为什么这个表达式是一个不错(尽管有偏)的 KL 估计器,以及如何在保持其低方差的同时使它无偏。
我们能计算 KL 的方法取决于我们对 和 的访问方式。这里我们假设可以计算任意 的概率(或概率密度) 和 ,但不能解析地对 求和(或积分)。为什么会不能解析求和呢?有可能精确计算需要过多的计算量或内存,没有解析解,或者我们只存储 log-prob 而不是整个分布,这能简化代码。这是合理的选择,尤其当 KL 只是作为诊断使用时(例如在强化学习中常见的情况)。最常见的求和或积分估计策略是 蒙特卡罗估计。给定样本 ,我们如何构造一个好的估计器?
一个好的估计器应当无偏(均值正确)且低方差。我们知道一个无偏估计器(PPO 的选择)是:
但它的方差很高,因为KL散度就定义而言一定是一个正值,但是对于上边的估计来说,它的一半的样本是负的(假设我们在没有任何对 和 的先验知识下),所以他的方差很高。为了简化表达,我们记:,那么原始的 KL 散度的估计可以写为:
为了减少方差,我们可以设计另外一个替代的 estimator:
它的方差低,但有偏。直觉上来说, 更好,因为每个样本都给出了 和 之间的“距离”,因而它总是正的。经验上来说, 的方差确实远低于 ,而且偏差也小了很多很多。至于为什么 比 的方差要低很多,原文中,利用f-散度给出了一个分析性的解释,在这里不再赘述。
那么有没有可能构造一个无偏并且方差又小的 KL 散度的估计器呢?一个降低方差的一般方法是使用控制变量(control variate):我们已知 是一个无偏估计,那么我们取 ,加上一个期望为零、且与其负相关的项,这样就能减小方差。我们能够轻而易举得到的一个期望为0的量是 。于是,对于任意 :
都是无偏的 KL 估计器。虽然理论上我们可以通过优化来最小化方差并解出最优 ,但结果依赖于 ,难以解析求出。注意到由于 是凹的:
因此若取 ,上式总为正。这里, 其实是 在 处的切线。所以当 , 我们其实是在测量 与其切线的垂直距离。所以我们可以得到估计器:
所以它总是正的。而 正是 GRPO 在 KL 散度估计上区别于 PPO 的一大特点 (PPO 使用的是 )。
在 RL for LLM 的角度讨论 KL 散度的估计
在 RL(比如 PPO、GRPO 等)里,我们经常会在 loss 里加一个 KL 散度,用来约束新策略(policy)不要偏离旧策略太多。这里, 是旧策略分布(,而 是新策略分布(), 则是一个完整的动作样本(在 LLM 里就是 token 或 token 序列)。我们通常用 表示 state(在 LLM 里就是 prompt 或上下文), 就是在这个上下文下生成的一个具体 token。在计算 KL 时,我们实际上是对 给定状态下的动作分布 做 KL,然后再对状态取期望:
在采样时,我们通常固定一个 prompt(state),然后针对它去估计这个 KL 值。
那么为什么无法直接计算 KL 的值,而一定要估计呢? 原因正如原文提到的三条,在 RL for LLM 里主要是 第 1 条:“动作空间(token space)太大,无法对所有可能的 枚举求和/积分。” 比如一个 tokenizer 有 50,000 个词表,即便只计算一个 token 的 KL,都要对这 50,000 个动作求和;而 RL 中通常是多步生成(sequence),这个空间就会是指数级的,因而根本不可行。另外,还有一个现实原因:在训练时,我们通常没有直接保存完整的分布(所有 token 的概率),只保存了生成路径上选中的 token 的 log-prob。这是为了节省显存和 I/O 开销。因此我们只能用 蒙特卡洛采样:从某个分布(通常是 ,即旧策略)采样一些 ,用这些样本来近似 KL。这里就回归到了上边这个博客所讨论的范畴。
在上边这篇文章里,我们一直在讨论的估计器(estimator),其实就是一个基于采样的函数,它输入的是在某个样本 上 和 的值(比如他们的比值 ),输出一个数,我们用这个数的样本均值用于近似 KL 散度。比如:
这些 就是不同的形式的 KL-estimator,它们都能通过在采样上取平均来近似 KL 散度,只是无偏性、方差大小不同。当我们选定了估计器,我们其实就决定用哪种公式去近似 KL。操作过程是:
- 采样 从旧策略 采样一批 token(或序列)。
- 计算 log-prob 对每个样本,计算新策略和旧策略下的对数概率:
并求出 或 。
- 代入 estimator 公式 比如选 ,就算:
- 求平均
这就是 KL 的近似值,用它替代真实 KL。
如果我们对比 KL 散度的真实计算过程(不估计),对于离散型概率分布(LLM 单步 token):我们要遍历每一个可能的token :
这里可以通过对比看出,使用 estimator 后计算量其实是巨大的。
关于不同 KL Estimator 的方差的讨论
需要指出的是,这里我们所说的方差,是说通过估计器采样出来的值的方差
即估计器 在样本空间上的波动程度。一个无偏的估计器,是指在无限多样本下,它的均值等于真实 KL;那么一个方差大的估计器,就可能造成即便期望正确(无偏),少量样本的平均值也可能偏离很远。在 RL for LLM 中,KL 项通常是 loss 里的一个正则化约束系数(比如 ),如果 KL 估计器的方差大,就会导致 loss 抖动,进而导致梯度抖动、训练不稳定。
原文中,作为为了让读者直观地理解为什么 不是一个小方差的估计器,给出了如下解释:
However, it () has high-variance, as it’s negative for half of the samples, whereas KL is always positive.
作者指出, 虽然是无偏估计,但是由于 和 在无先验知识的条件下,采样出来的样本值总是一半比对方大,一半比对方小,所以导致 采集到的数据总是一部分大于 另一半小于 ,到这里我认为都是没问题的,但是作者后边说:因为 KL 总是大于 0 的 (可由基本不等式推出),所以 的方差很高。这里因果关系不成立的原因在于,我们不能用期望值的正负去要求估计值采样的正负。一个简单的反驳是,如果计算期望, 也是一半是正一半是负,这并不能说明什么。其实单次采样的 log 比值(无论是 还是 )确实都可能正也可能负,和 一样会“跳变”,所以正负跳变本身并不是导致方差大的唯一原因。
根据 KL 散度的定义:
这个期望是必然非负,但里面的被积项 在单个样本上可以正也可以负。 其实就是直接取这个被积项:
所以它的每个样本值确实可能正也可能负,和 KL 定义中的 integrand 一样。
那为什么 方差大?
不是因为“正负跳变”本身,而是因为 的取值分布通常很宽(heavy-tailed),例如,在概率很小的样本 下, 可以变得非常大(正或负)。这种极端值在有限样本平均中占据很大权重,所以会导致高方差。换句话说,问题是极端值 + 正负抵消的组合:正负抵消让你需要更多样本才能收敛到真实均值,极端值让样本方差本身更大。所以文章里提的“因为一半是负的”只是一个直观类比,并不是完整的方差来源解释。
如果我们从这个角度,来看其他估计器 和 ,就会发现: 永远为正,去掉了正负抵消的可能,因为带来了偏差;但是引入了平方,使得数值变得更平滑,使方差下降; 通过控制变量消掉了一部分波动源,在保留无偏性的同时降低了方差,具体请看下边的解释。
在 PPO/GRPO 里,如果用 ,当 batch 小、分布差距大时,KL 估计值会很抖动(因为几个极端样本能显著改变均值),同时会使 KL-penalty 系数调节不稳(可能突然过强或过弱),换成低方差估计器( 或 )时,batch 中的每个样本对 KL 的贡献更稳定,不容易被少数极端样本主导。
为什么 能做到无偏且小方差?
乍一看, 全是正数,直觉会以为它的均值要大于 的均值。
但别忘了 是通过 控制变量(control variate) 从 推出来的。文章中给的思路是:
其中 ,在 下的期望是:
所以加上任何倍数的 ,不会改变期望值。当取 时:
这就解释了为什么 的期望等于 的期望,也等于 KL,即是一个无偏估计。
比 的方差小的原因在于: 里只有 ,波动可能很大(正负都有,大值出现时特别大);而 和 在数值上高度相关(一个变大时另一个也变大/变小),而且相关性是负的,加上 相当于用一个负相关的量去抵消波动。抵消之后剩下的 ,取值范围收紧,且全为正值,最终会导致样本方差下降。
@misc{zhang2025kl,
author = {Zhang, Yihua},
title = {从 RL for LLM 视角重新理解 KL 估计:读《Approximating KL Divergence》笔记},
year = {2025},
url = {https://normaluhr.github.io/zh/p/kl}
}