on-policy 采样的浪费
策略梯度 那章的结尾,给出了一个可用的估计量:
∇ θ J ( θ ) = E τ ∼ p θ [ ∑ t = 1 T ∇ θ log π θ ( a t ∣ s t ) A ^ t ] \nabla_\theta J(\theta)
=
\mathbb{E}_{\tau \sim p_\theta}
\left[
\sum_{t=1}^{T}
\nabla_\theta \log \pi_\theta(a_t \mid s_t)\, \widehat A_t
\right] ∇ θ J ( θ ) = E τ ∼ p θ [ t = 1 ∑ T ∇ θ log π θ ( a t ∣ s t ) A t ]
本章的全部内容,都藏在这一行的一个细节里:那个下标。期望是对从 p θ p_\theta p θ 采样出的轨迹取的,而 p θ p_\theta p θ 是由当前 参数诱导出的分布。只要走一步梯度,θ \theta θ 就不再是生成这批数据的那个 θ \theta θ 了。手里的每一个样本此刻都来自错误的分布,而那个促成了这一步的估计量,没有任何依据去促成第二步。
于是这个循环是被逼出来的:采一批、走一步、把这批扔掉、再采一批。每条轨迹只支撑一次更新,然后就被丢弃。这里没有任何东西被算了两遍,也没有任何东西被浪费了两次;浪费在于,昂贵的样本只被用了一次。
有多昂贵取决于问题,而对本系列正在走向的那些领域,答案是:昂贵到足以盖过其他一切。采一个样本意味着生成一整条 token 序列,或者跑完一整条去噪链。相比之下梯度那一步是廉价的。花掉一整轮生成,只换来一步更新,这就是 on-policy 策略梯度的核心低效,也正是近端策略优化存在的理由。
所以本章要回答的问题既窄又实际:
一批轨迹能不能支撑多次梯度更新?
答案是可以,但有条件,而本章余下的部分讲的就是这个条件要付什么代价。
我们在估什么
在动采样这套机器之前,先把那个要与梯度相乘的对象讲精确是值得的,因为下面每一步都在摆弄它,而最终的目标函数会按它的符号分情形。
定义 Advantage 函数
对策略 π \pi π ,advantage 函数是
A π ( s , a ) = Q π ( s , a ) − V π ( s ) . A^\pi(s,a) = Q^\pi(s,a) - V^\pi(s). A π ( s , a ) = Q π ( s , a ) − V π ( s ) .
两项都来自回报、价值与 Bellman 方程 :Q π ( s , a ) Q^\pi(s,a) Q π ( s , a ) 是在 s s s 处采取 a a a 、其后遵循 π \pi π 的期望回报,V π ( s ) V^\pi(s) V π ( s ) 是从 s s s 出发、连第一个动作也由 π \pi π 采样时的期望回报。两者之差,就是”特意采取 a a a “相对于”让策略自己拿主意”多买到的那部分。
这个”相对于”不是随口一说。下面这条命题说的就是它。
命题 Advantage 在自己的策略下均值为零
对每一个状态 s s s ,
E a ∼ π ( ⋅ ∣ s ) [ A π ( s , a ) ] = 0. \mathbb{E}_{a \sim \pi(\cdot \mid s)}\big[A^\pi(s,a)\big] = 0. E a ∼ π ( ⋅ ∣ s ) [ A π ( s , a ) ] = 0.
证明
展开定义,把期望拆成两项:
E a ∼ π ( ⋅ ∣ s ) [ A π ( s , a ) ] = ∑ a π ( a ∣ s ) [ Q π ( s , a ) − V π ( s ) ] = ∑ a π ( a ∣ s ) Q π ( s , a ) − V π ( s ) ∑ a π ( a ∣ s ) . \mathbb{E}_{a \sim \pi(\cdot\mid s)}\big[A^\pi(s,a)\big]
=
\sum_a \pi(a\mid s)\big[Q^\pi(s,a) - V^\pi(s)\big]
=
\sum_a \pi(a\mid s)Q^\pi(s,a)
-
V^\pi(s)\sum_a \pi(a\mid s). E a ∼ π ( ⋅ ∣ s ) [ A π ( s , a ) ] = a ∑ π ( a ∣ s ) [ Q π ( s , a ) − V π ( s ) ] = a ∑ π ( a ∣ s ) Q π ( s , a ) − V π ( s ) a ∑ π ( a ∣ s ) . 第一个求和由前面章节”由动作价值得到状态价值”那条命题,恰好是 V π ( s ) V^\pi(s) V π ( s ) 。第二个求和是 1 1 1 ,因为 π ( ⋅ ∣ s ) \pi(\cdot\mid s) π ( ⋅ ∣ s ) 是一个概率分布。于是两项都等于 V π ( s ) V^\pi(s) V π ( s ) ,相消。
这就把符号的含义钉死了,而符号恰恰是 clip 目标函数唯一真正用到的性质:
A π ( s , a ) > 0 A^\pi(s,a) > 0 A π ( s , a ) > 0 :采取 a a a 比 π \pi π 在 s s s 处的平均表现更好。更新应该让 a a a 更可能被采到。
A π ( s , a ) < 0 A^\pi(s,a) < 0 A π ( s , a ) < 0 :采取 a a a 比 π \pi π 自己的平均还差。更新应该让 a a a 更不可能被采到。
因为均值恰好是零,“比平均更好”是一句 π \pi π 拿自己当尺子的话。advantage 不是对一个动作的孤立评分,它是对照着正在被改进的那个策略 做的比较,而且那个策略一变,它就得重算。最后这半句值得记着,它是这批样本会变馊的第二个理由。
实践中 A π A^\pi A π 是未知的,必须从这批数据里估计。策略梯度 那章搭好了本章要继承的估计量 A ^ t = G t − b ( s t ) \widehat A_t = G_t - b(s_t) A t = G t − b ( s t ) ,取 b ( s t ) = V π ( s t ) b(s_t) = V^\pi(s_t) b ( s t ) = V π ( s t ) ,并在那里证明了:减去任何与动作无关的 baseline,都不会给梯度方向引入偏差。下面的内容不依赖 A ^ t \widehat A_t A t 具体怎么来,只依赖两件事:它估计的是 A π θ old A^{\pi_{\theta_{\text{old}}}} A π θ old ,以及它带着一个符号。
重要性采样
障碍在于,期望是对错误的分布取的。而恰好有一个标准恒等式,就是为这种情形准备的。
引理 重要性采样恒等式
设 p p p 和 q q q 是同一空间上的两个分布,f f f 是一个函数,且对每一个满足 p ( x ) f ( x ) ≠ 0 p(x)f(x) \neq 0 p ( x ) f ( x ) = 0 的 x x x 都有 q ( x ) > 0 q(x) > 0 q ( x ) > 0 。那么
E x ∼ p [ f ( x ) ] = E x ∼ q [ p ( x ) q ( x ) f ( x ) ] . \mathbb{E}_{x \sim p}\big[f(x)\big]
=
\mathbb{E}_{x \sim q}\left[\frac{p(x)}{q(x)}f(x)\right]. E x ∼ p [ f ( x ) ] = E x ∼ q [ q ( x ) p ( x ) f ( x ) ] .
证明
把左边写成求和,再插入 q ( x ) / q ( x ) q(x)/q(x) q ( x ) / q ( x ) 。由支撑集条件,这个操作在每一个贡献非零项的 x x x 处都是合法的:
E x ∼ p [ f ( x ) ] = ∑ x p ( x ) f ( x ) = ∑ x q ( x ) p ( x ) q ( x ) f ( x ) = E x ∼ q [ p ( x ) q ( x ) f ( x ) ] . \mathbb{E}_{x\sim p}[f(x)]
=
\sum_x p(x)f(x)
=
\sum_x q(x)\,\frac{p(x)}{q(x)}\,f(x)
=
\mathbb{E}_{x\sim q}\left[\frac{p(x)}{q(x)}f(x)\right]. E x ∼ p [ f ( x )] = x ∑ p ( x ) f ( x ) = x ∑ q ( x ) q ( x ) p ( x ) f ( x ) = E x ∼ q [ q ( x ) p ( x ) f ( x ) ] . 满足 p ( x ) f ( x ) = 0 p(x)f(x) = 0 p ( x ) f ( x ) = 0 的那些点,对两边的贡献都是零,因此无论 q ( x ) q(x) q ( x ) 取何值都可以忽略。x x x 连续时求和换成积分,论证不变。
把这个恒等式用到轨迹上,会暴露出让这一切可计算的那份运气。
命题 轨迹比值中环境项相消
p θ ( τ ) p θ old ( τ ) = ∏ t = 1 T π θ ( a t ∣ s t ) π θ old ( a t ∣ s t ) \frac{p_\theta(\tau)}{p_{\theta_{\text{old}}}(\tau)}
=
\prod_{t=1}^{T}
\frac{\pi_\theta(a_t \mid s_t)}
{\pi_{\theta_{\text{old}}}(a_t \mid s_t)} p θ old ( τ ) p θ ( τ ) = t = 1 ∏ T π θ old ( a t ∣ s t ) π θ ( a t ∣ s t )
证明
上一章的轨迹分解给出
p θ ( τ ) = p ( s 1 ) ∏ t = 1 T π θ ( a t ∣ s t ) p ( s t + 1 ∣ s t , a t ) , p_\theta(\tau)
=
p(s_1)\prod_{t=1}^{T}\pi_\theta(a_t\mid s_t)\,p(s_{t+1}\mid s_t,a_t), p θ ( τ ) = p ( s 1 ) t = 1 ∏ T π θ ( a t ∣ s t ) p ( s t + 1 ∣ s t , a t ) , 以及把 θ \theta θ 换成 θ old \theta_{\text{old}} θ old 的同一个式子。初始状态项 p ( s 1 ) p(s_1) p ( s 1 ) 和每一个转移项 p ( s t + 1 ∣ s t , a t ) p(s_{t+1}\mid s_t,a_t) p ( s t + 1 ∣ s t , a t ) 都来自环境,不带任何对策略参数的依赖,因此它们在分子分母里一模一样地出现,逐项相消,只剩下策略项。
这和”让策略梯度不需要模型就能计算”的那次相消是同一个机制,只是第二次出现、出于不同的理由。那里它消掉的是 ∇ θ \nabla_\theta ∇ θ 作用在智能体控制不了的项上;这里它把那些项从一个比值里消掉。两次的结论是一样的:一个看上去需要知道环境动态的量,其实不需要。
定义 概率比值
r t ( θ ) = π θ ( a t ∣ s t ) π θ old ( a t ∣ s t ) r_t(\theta)
=
\frac{\pi_\theta(a_t \mid s_t)}
{\pi_{\theta_{\text{old}}}(a_t \mid s_t)} r t ( θ ) = π θ old ( a t ∣ s t ) π θ ( a t ∣ s t )
定义 Surrogate 目标
L ( θ ) = E t [ r t ( θ ) A ^ t ] L(\theta)
=
\mathbb{E}_{t}\big[r_t(\theta)\,\widehat A_t\big] L ( θ ) = E t [ r t ( θ ) A t ] 其中期望是对用 π θ old \pi_{\theta_{\text{old}}} π θ old 采样出的轨迹的各时间步取的,A ^ t \widehat A_t A t 估计 A π θ old ( s t , a t ) A^{\pi_{\theta_{\text{old}}}}(s_t,a_t) A π θ old ( s t , a t ) 。
L L L 之所以值得优化,是因为在采集样本的那一点上,它的梯度恰好就是上一章推出来的那个对象。
命题 Surrogate 在采样点处一阶正确
∇ θ L ( θ ) ∣ θ = θ old = ∇ θ J ( θ ) ∣ θ = θ old \nabla_\theta L(\theta)\big|_{\theta = \theta_{\text{old}}}
=
\nabla_\theta J(\theta)\big|_{\theta = \theta_{\text{old}}} ∇ θ L ( θ ) θ = θ old = ∇ θ J ( θ ) θ = θ old
证明
对比值求导,并在 θ = θ old \theta = \theta_{\text{old}} θ = θ old 处取值。分母不依赖 θ \theta θ ,所以
∇ θ r t ( θ ) ∣ θ old = ∇ θ π θ ( a t ∣ s t ) ∣ θ old π θ old ( a t ∣ s t ) = ∇ θ log π θ ( a t ∣ s t ) ∣ θ old , \nabla_\theta r_t(\theta)\Big|_{\theta_{\text{old}}}
=
\frac{\nabla_\theta \pi_\theta(a_t\mid s_t)\big|_{\theta_{\text{old}}}}
{\pi_{\theta_{\text{old}}}(a_t\mid s_t)}
=
\nabla_\theta \log \pi_\theta(a_t\mid s_t)\Big|_{\theta_{\text{old}}}, ∇ θ r t ( θ ) θ old = π θ old ( a t ∣ s t ) ∇ θ π θ ( a t ∣ s t ) θ old = ∇ θ log π θ ( a t ∣ s t ) θ old , 最后一步是上一章那个 log-derivative 恒等式 ∇ log π = ∇ π / π \nabla \log \pi = \nabla \pi / \pi ∇ log π = ∇ π / π ,从右往左读。又因为 A ^ t \widehat A_t A t 不依赖 θ \theta θ ,
∇ θ L ( θ ) ∣ θ old = E t [ ∇ θ log π θ ( a t ∣ s t ) A ^ t ] ∣ θ old , \nabla_\theta L(\theta)\big|_{\theta_{\text{old}}}
=
\mathbb{E}_t\big[\nabla_\theta \log \pi_\theta(a_t\mid s_t)\,\widehat A_t\big]\Big|_{\theta_{\text{old}}}, ∇ θ L ( θ ) θ old = E t [ ∇ θ log π θ ( a t ∣ s t ) A t ] θ old , 这恰好就是策略梯度估计量。
这条命题很容易被读过头。它没有说 L = J L = J L = J ,也没有说 ∇ L = ∇ J \nabla L = \nabla J ∇ L = ∇ J 。它说的是这两个梯度在某一个点 上相等,即 θ old \theta_{\text{old}} θ old ,而对其他任何 θ \theta θ 它什么都没说。
这个区别是整章的枢纽。L L L 是 J J J 的一个局部替身,在这批数据被采出的那个位置上一阶精确,而它的有效期随着 θ \theta θ 的移动而过期。把 L L L 当成”要最大化的那个东西”,而不是当成”一个在 θ old \theta_{\text{old}} θ old 附近可信、离远了就无意义的模型”,正是本章余下每一个机制所要防止的错误。
无偏还不够
重要性采样恒等式是精确的。于是很容易得出”重用是免费的”这个结论,而这个结论是错的。恒等式等同的是两个期望 ,而期望是关于无穷多个样本的陈述。一批有限的数据能交出什么,取决于方差。
命题 重要性采样估计量的方差
在重要性采样引理的条件下,
Var x ∼ q [ p ( x ) q ( x ) f ( x ) ] = E x ∼ p [ p ( x ) q ( x ) f ( x ) 2 ] − ( E x ∼ p [ f ( x ) ] ) 2 , \operatorname{Var}_{x\sim q}\left[\frac{p(x)}{q(x)}f(x)\right]
=
\mathbb{E}_{x\sim p}\left[\frac{p(x)}{q(x)}f(x)^2\right]
-
\big(\mathbb{E}_{x\sim p}[f(x)]\big)^2, Var x ∼ q [ q ( x ) p ( x ) f ( x ) ] = E x ∼ p [ q ( x ) p ( x ) f ( x ) 2 ] − ( E x ∼ p [ f ( x )] ) 2 , 将它与下式对照:
Var x ∼ p [ f ( x ) ] = E x ∼ p [ f ( x ) 2 ] − ( E x ∼ p [ f ( x ) ] ) 2 . \operatorname{Var}_{x\sim p}\big[f(x)\big]
=
\mathbb{E}_{x\sim p}\big[f(x)^2\big]
-
\big(\mathbb{E}_{x\sim p}[f(x)]\big)^2 . Var x ∼ p [ f ( x ) ] = E x ∼ p [ f ( x ) 2 ] − ( E x ∼ p [ f ( x )] ) 2 .
证明
对 q q q 下的 X = ( p / q ) f X = (p/q)f X = ( p / q ) f 应用 Var [ X ] = E [ X 2 ] − ( E [ X ] ) 2 \operatorname{Var}[X] = \mathbb{E}[X^2] - (\mathbb{E}[X])^2 Var [ X ] = E [ X 2 ] − ( E [ X ] ) 2 。第二项由引理等于 ( E x ∼ q [ ( p / q ) f ] ) 2 = ( E x ∼ p [ f ] ) 2 (\mathbb{E}_{x\sim q}[(p/q)f])^2 = (\mathbb{E}_{x\sim p}[f])^2 ( E x ∼ q [( p / q ) f ] ) 2 = ( E x ∼ p [ f ] ) 2 。至于第一项,
E x ∼ q [ ( p ( x ) q ( x ) f ( x ) ) 2 ] = ∑ x q ( x ) p ( x ) 2 q ( x ) 2 f ( x ) 2 = ∑ x p ( x ) p ( x ) q ( x ) f ( x ) 2 = E x ∼ p [ p ( x ) q ( x ) f ( x ) 2 ] . \mathbb{E}_{x\sim q}\left[\left(\frac{p(x)}{q(x)}f(x)\right)^2\right]
=
\sum_x q(x)\frac{p(x)^2}{q(x)^2}f(x)^2
=
\sum_x p(x)\,\frac{p(x)}{q(x)}\,f(x)^2
=
\mathbb{E}_{x\sim p}\left[\frac{p(x)}{q(x)}f(x)^2\right]. E x ∼ q [ ( q ( x ) p ( x ) f ( x ) ) 2 ] = x ∑ q ( x ) q ( x ) 2 p ( x ) 2 f ( x ) 2 = x ∑ p ( x ) q ( x ) p ( x ) f ( x ) 2 = E x ∼ p [ q ( x ) p ( x ) f ( x ) 2 ] .
两个式子只在一个地方不同:重加权后的估计量在二阶矩里多带了一个因子 p ( x ) / q ( x ) p(x)/q(x) p ( x ) / q ( x ) 。凡是新策略比旧策略放了明显更多质量的地方,这个因子就很大,而它是乘性地 进入方差,同时对均值毫无贡献。
这是第一个失效。第二个已经摆在桌面上了,而且更根本。
上一节那条命题是在 θ old \theta_{\text{old}} θ old 处 成立的。θ \theta θ 一旦离开,关于 L L L 与 J J J 的关系就什么都没被确立过。把 θ \theta θ 推得足够远,L L L 上升不蕴含 J J J 的任何事情:优化器正在勤勤恳恳地爬一个模型,而这个模型唯一的保证是局部的,且它已经把那个局部抛在身后了。上面标出的那个逐时间步近似,按同样的时刻表、出于同样的理由一起退化。
两个互相独立的失效,指向同一个结论。
重加权估计量的方差,随 π θ \pi_\theta π θ 与 π θ old \pi_{\theta_{\text{old}}} π θ old 的分离而增长。surrogate 作为 J J J 的模型,其有效性只在 θ old \theta_{\text{old}} θ old 处被保证,并随 θ \theta θ 离开而退化。
重用是合法的,但只在局部合法。 本章余下的每一节,都是这句话的推论:先把”局部”讲精确,再找一个廉价的办法维持它。
多近才算近:KL 散度
“待在 θ old \theta_{\text{old}} θ old 附近”还不是一条可用的指令,因为它没说在什么意义上的附近。最朴素的读法(在参数空间里靠近)恰恰是错的:∥ θ − θ old ∥ \lVert \theta - \theta_{\text{old}} \rVert ∥ θ − θ old ∥ 与”策略的行为 改变了多少”之间没有固定关系。同样大小的参数位移,在参数空间的某个区域可能让策略几乎纹丝不动,在另一个区域却足以把它彻底改写。上一节那两个失效真正依赖的,是两个分布 之间的分离,所以要量的就是它。
定义 KL 散度
对同一空间上的分布 p p p 和 q q q ,且在所有 p ( x ) > 0 p(x) > 0 p ( x ) > 0 处都有 q ( x ) > 0 q(x) > 0 q ( x ) > 0 ,
K L ( p ∥ q ) = E x ∼ p [ log p ( x ) q ( x ) ] = ∑ x p ( x ) log p ( x ) q ( x ) . \mathrm{KL}(p \,\Vert\, q)
=
\mathbb{E}_{x \sim p}\left[\log \frac{p(x)}{q(x)}\right]
=
\sum_x p(x) \log \frac{p(x)}{q(x)}. KL ( p ∥ q ) = E x ∼ p [ log q ( x ) p ( x ) ] = x ∑ p ( x ) log q ( x ) p ( x ) .
命题 非负性
K L ( p ∥ q ) ≥ 0 \mathrm{KL}(p \Vert q) \ge 0 KL ( p ∥ q ) ≥ 0 ,取等当且仅当 p = q p = q p = q (在 p p p 下几乎处处)。
证明
对凹函数 log \log log 应用 Jensen 不等式:
− K L ( p ∥ q ) = E x ∼ p [ log q ( x ) p ( x ) ] ≤ log E x ∼ p [ q ( x ) p ( x ) ] = log ∑ x p ( x ) q ( x ) p ( x ) = log ∑ x q ( x ) = log 1 = 0 , -\mathrm{KL}(p\Vert q)
=
\mathbb{E}_{x\sim p}\left[\log \frac{q(x)}{p(x)}\right]
\le
\log \mathbb{E}_{x\sim p}\left[\frac{q(x)}{p(x)}\right]
=
\log \sum_x p(x)\frac{q(x)}{p(x)}
=
\log \sum_x q(x)
=
\log 1
=
0, − KL ( p ∥ q ) = E x ∼ p [ log p ( x ) q ( x ) ] ≤ log E x ∼ p [ p ( x ) q ( x ) ] = log x ∑ p ( x ) p ( x ) q ( x ) = log x ∑ q ( x ) = log 1 = 0 , 于是 K L ( p ∥ q ) ≥ 0 \mathrm{KL}(p\Vert q) \ge 0 KL ( p ∥ q ) ≥ 0 。又因为 log \log log 是严格 凹的,Jensen 取等当且仅当 q ( x ) / p ( x ) q(x)/p(x) q ( x ) / p ( x ) 在 p p p 下几乎处处为常数;而 p p p 与 q q q 都求和为一,这个常数只能是 1 1 1 ,即 p = q p = q p = q 几乎处处。
有两个事实把它和比值连了起来,而它们其实是同一个事实。
命题 比值的平均恒为一
对任意状态 s s s ,期望对旧策略自己的动作取,
E a ∼ π θ old ( ⋅ ∣ s ) [ r ( θ ) ] = 1. \mathbb{E}_{a \sim \pi_{\theta_{\text{old}}}(\cdot\mid s)}\big[r(\theta)\big] = 1 . E a ∼ π θ old ( ⋅ ∣ s ) [ r ( θ ) ] = 1.
证明
E a ∼ π θ old ( ⋅ ∣ s ) [ π θ ( a ∣ s ) π θ old ( a ∣ s ) ] = ∑ a π θ old ( a ∣ s ) π θ ( a ∣ s ) π θ old ( a ∣ s ) = ∑ a π θ ( a ∣ s ) = 1 , \mathbb{E}_{a\sim\pi_{\theta_{\text{old}}}(\cdot\mid s)}\left[\frac{\pi_\theta(a\mid s)}{\pi_{\theta_{\text{old}}}(a\mid s)}\right]
=
\sum_a \pi_{\theta_{\text{old}}}(a\mid s)\,\frac{\pi_\theta(a\mid s)}{\pi_{\theta_{\text{old}}}(a\mid s)}
=
\sum_a \pi_\theta(a\mid s)
=
1, E a ∼ π θ old ( ⋅ ∣ s ) [ π θ old ( a ∣ s ) π θ ( a ∣ s ) ] = a ∑ π θ old ( a ∣ s ) π θ old ( a ∣ s ) π θ ( a ∣ s ) = a ∑ π θ ( a ∣ s ) = 1 , 其中用到支撑集条件来相消,以及 π θ ( ⋅ ∣ s ) \pi_\theta(\cdot\mid s) π θ ( ⋅ ∣ s ) 是一个分布。
命题 KL 就是负对数比值的期望
K L ( π θ old ( ⋅ ∣ s ) ∥ π θ ( ⋅ ∣ s ) ) = E a ∼ π θ old ( ⋅ ∣ s ) [ − log r ( θ ) ] \mathrm{KL}\big(\pi_{\theta_{\text{old}}}(\cdot\mid s) \,\Vert\, \pi_\theta(\cdot\mid s)\big)
=
\mathbb{E}_{a \sim \pi_{\theta_{\text{old}}}(\cdot\mid s)}\big[-\log r(\theta)\big] KL ( π θ old ( ⋅ ∣ s ) ∥ π θ ( ⋅ ∣ s ) ) = E a ∼ π θ old ( ⋅ ∣ s ) [ − log r ( θ ) ]
证明
直接由定义:
K L ( π θ old ∥ π θ ) = E a ∼ π θ old [ log π θ old ( a ∣ s ) π θ ( a ∣ s ) ] = E a ∼ π θ old [ − log π θ ( a ∣ s ) π θ old ( a ∣ s ) ] = E a ∼ π θ old [ − log r ( θ ) ] . \mathrm{KL}\big(\pi_{\theta_{\text{old}}} \Vert \pi_\theta\big)
=
\mathbb{E}_{a\sim\pi_{\theta_{\text{old}}}}\left[\log \frac{\pi_{\theta_{\text{old}}}(a\mid s)}{\pi_\theta(a\mid s)}\right]
=
\mathbb{E}_{a\sim\pi_{\theta_{\text{old}}}}\left[-\log \frac{\pi_\theta(a\mid s)}{\pi_{\theta_{\text{old}}}(a\mid s)}\right]
=
\mathbb{E}_{a\sim\pi_{\theta_{\text{old}}}}\big[-\log r(\theta)\big]. KL ( π θ old ∥ π θ ) = E a ∼ π θ old [ log π θ ( a ∣ s ) π θ old ( a ∣ s ) ] = E a ∼ π θ old [ − log π θ old ( a ∣ s ) π θ ( a ∣ s ) ] = E a ∼ π θ old [ − log r ( θ ) ] .
推论 非负性,再证一次,这次用本章自己的记号
把上面两条命题合起来,并用 − log -\log − log 是凸的:
K L ( π θ old ∥ π θ ) = E [ − log r ] ≥ − log E [ r ] = − log 1 = 0. \mathrm{KL}\big(\pi_{\theta_{\text{old}}} \Vert \pi_\theta\big)
=
\mathbb{E}\big[-\log r\big]
\ \ge\
-\log \mathbb{E}\big[r\big]
=
-\log 1
=
0 . KL ( π θ old ∥ π θ ) = E [ − log r ] ≥ − log E [ r ] = − log 1 = 0.
这条推论值得停一下,因为它说明:被控制的那个散度,和被 clip 的那个比值,不是两个话题。比值在旧策略下平均恰好为一;散度就是同一个比值取负对数之后的期望;而 Jensen 不等式在 E [ − log r ] \mathbb{E}[-\log r] E [ − log r ] 与 − log E [ r ] -\log\mathbb{E}[r] − log E [ r ] 之间留下的那道缝,就是 这个散度。一个没有移动过的策略,r ≡ 1 r \equiv 1 r ≡ 1 ,散度为零。一个移动过的策略,比值散布在一附近,而散度量的就是这份散布在对数尺度上的大小。本章对 r r r 做的一切,都在间接地对 K L \mathrm{KL} KL 做。
Clipping
这副药必须在维持”待在近处”的同时,仍然是一阶方法:一个梯度,不要二阶导,不要带约束的优化,在一个普通训练循环里几行代码就能写完。近端策略优化的答案是:不动优化器,改目标函数,让”走远”这件事干脆不再划算。
定义 Clipped Surrogate 目标
L C L I P ( θ ) = E t [ min ( r t ( θ ) A ^ t , clip ( r t ( θ ) , 1 − ϵ , 1 + ϵ ) A ^ t ) ] L^{\mathrm{CLIP}}(\theta)
=
\mathbb{E}_t
\Big[
\min\big(
r_t(\theta)\,\widehat A_t,
\ \operatorname{clip}\big(r_t(\theta), 1-\epsilon, 1+\epsilon\big)\,\widehat A_t
\big)
\Big] L CLIP ( θ ) = E t [ min ( r t ( θ ) A t , clip ( r t ( θ ) , 1 − ϵ , 1 + ϵ ) A t ) ] 其中 clip ( r , a , b ) = max ( a , min ( r , b ) ) \operatorname{clip}(r, a, b) = \max(a, \min(r, b)) clip ( r , a , b ) = max ( a , min ( r , b )) ,ϵ \epsilon ϵ 是一个小常数,典型取 0.2 0.2 0.2 。
这个定义很紧凑,而它的行为从字面上是读不出来的。下面这条命题把它讲明。
命题 clip 是一个单侧的梯度掩码
对单个样本记 L t ( θ ) = min ( r t A ^ t , clip ( r t , 1 − ϵ , 1 + ϵ ) A ^ t ) L_t(\theta) = \min\big(r_t\widehat A_t, \operatorname{clip}(r_t, 1-\epsilon, 1+\epsilon)\widehat A_t\big) L t ( θ ) = min ( r t A t , clip ( r t , 1 − ϵ , 1 + ϵ ) A t ) ,并把它看作 r t r_t r t 的函数。那么在导数存在的每一处,
∂ L t ∂ r t = { 0 若 A ^ t > 0 且 r t > 1 + ϵ , 0 若 A ^ t < 0 且 r t < 1 − ϵ , A ^ t 其余情形。 \frac{\partial L_t}{\partial r_t}
=
\begin{cases}
0 & \text{若 } \widehat A_t > 0 \text{ 且 } r_t > 1+\epsilon,\\
0 & \text{若 } \widehat A_t < 0 \text{ 且 } r_t < 1-\epsilon,\\
\widehat A_t & \text{其余情形。}
\end{cases} ∂ r t ∂ L t = ⎩ ⎨ ⎧ 0 0 A t 若 A t > 0 且 r t > 1 + ϵ , 若 A t < 0 且 r t < 1 − ϵ , 其余情形。
证明
依次考虑 A ^ t \widehat A_t A t 的两种符号,并在每种符号下考虑 r t r_t r t 的三个区域。
设 A ^ t > 0 \widehat A_t > 0 A t > 0 。当 r t > 1 + ϵ r_t > 1+\epsilon r t > 1 + ϵ ,clip 返回 1 + ϵ 1+\epsilon 1 + ϵ ,而乘以一个正的 A ^ t \widehat A_t A t 保序,故 r t A ^ t > ( 1 + ϵ ) A ^ t r_t \widehat A_t > (1+\epsilon)\widehat A_t r t A t > ( 1 + ϵ ) A t ,min \min min 选中 ( 1 + ϵ ) A ^ t (1+\epsilon)\widehat A_t ( 1 + ϵ ) A t ,它对 r t r_t r t 是常数,导数为 0 0 0 。当 r t ∈ [ 1 − ϵ , 1 + ϵ ] r_t \in [1-\epsilon, 1+\epsilon] r t ∈ [ 1 − ϵ , 1 + ϵ ] ,clip 返回 r t r_t r t ,两个参数都是 r t A ^ t r_t\widehat A_t r t A t ,导数为 A ^ t \widehat A_t A t 。当 r t < 1 − ϵ r_t < 1-\epsilon r t < 1 − ϵ ,clip 返回 1 − ϵ 1-\epsilon 1 − ϵ ,而 r t A ^ t < ( 1 − ϵ ) A ^ t r_t \widehat A_t < (1-\epsilon)\widehat A_t r t A t < ( 1 − ϵ ) A t ,min \min min 选中未 clip 的 r t A ^ t r_t\widehat A_t r t A t ,导数为 A ^ t \widehat A_t A t 。
设 A ^ t < 0 \widehat A_t < 0 A t < 0 。乘以一个负数会反序,于是选择结果对调。当 r t < 1 − ϵ r_t < 1-\epsilon r t < 1 − ϵ ,r t A ^ t > ( 1 − ϵ ) A ^ t r_t\widehat A_t > (1-\epsilon)\widehat A_t r t A t > ( 1 − ϵ ) A t ,min \min min 选中 ( 1 − ϵ ) A ^ t (1-\epsilon)\widehat A_t ( 1 − ϵ ) A t ,常数,导数为 0 0 0 。当 r t ∈ [ 1 − ϵ , 1 + ϵ ] r_t \in [1-\epsilon,1+\epsilon] r t ∈ [ 1 − ϵ , 1 + ϵ ] ,两个参数都是 r t A ^ t r_t\widehat A_t r t A t ,导数为 A ^ t \widehat A_t A t 。当 r t > 1 + ϵ r_t > 1+\epsilon r t > 1 + ϵ ,r t A ^ t < ( 1 + ϵ ) A ^ t r_t\widehat A_t < (1+\epsilon)\widehat A_t r t A t < ( 1 + ϵ ) A t ,min \min min 选中未 clip 的 r t A ^ t r_t\widehat A_t r t A t ,导数为 A ^ t \widehat A_t A t 。
把六种情形收拢即得结论;导数不存在的地方只有两个折点 r t ∈ { 1 − ϵ , 1 + ϵ } r_t \in \{1-\epsilon, 1+\epsilon\} r t ∈ { 1 − ϵ , 1 + ϵ } 。
摆成一张表,这个结构比从代数里读出来要容易看:
r t < 1 − ϵ r_t < 1-\epsilon r t < 1 − ϵ r t ∈ [ 1 − ϵ , 1 + ϵ ] r_t \in [1-\epsilon,\, 1+\epsilon] r t ∈ [ 1 − ϵ , 1 + ϵ ] r t > 1 + ϵ r_t > 1+\epsilon r t > 1 + ϵ A ^ t > 0 \widehat A_t > 0 A t > 0 A ^ t \widehat A_t A t A ^ t \widehat A_t A t 0 \mathbf{0} 0 A ^ t < 0 \widehat A_t < 0 A t < 0 0 \mathbf{0} 0 A ^ t \widehat A_t A t A ^ t \widehat A_t A t
六格里恰好两格梯度消失,而它们是同一格换了身衣服。A ^ t > 0 \widehat A_t > 0 A t > 0 时,更新想抬高 π θ ( a t ∣ s t ) \pi_\theta(a_t\mid s_t) π θ ( a t ∣ s t ) ,也就是抬高 r t r_t r t ;一旦 r t r_t r t 被抬过 1 + ϵ 1+\epsilon 1 + ϵ ,梯度就关掉。A ^ t < 0 \widehat A_t < 0 A t < 0 时,更新想压低概率,也就是压低 r t r_t r t ;一旦 r t r_t r t 跌破 1 − ϵ 1-\epsilon 1 − ϵ ,梯度就关掉。两种情形里,一旦这个样本已经朝它想去的方向被挪动了超过 ϵ \epsilon ϵ 的幅度,目标函数就不再为继续挪动付钱。
另外四格才是理解这东西为什么能成立的关键。
算法
一切装配起来,就是一个比 REINFORCE 多一层的循环:一个内层循环,把同一批数据反复花掉。
算法 1 使用 clip 目标的 PPO
Require: 初始策略参数 θ \theta θ ,clip 宽度 ϵ \epsilon ϵ ,epoch 数 K K K ,步长 α \alpha α
1: repeat
2: θ old ← θ \theta_{\text{old}} \gets \theta θ old ← θ
3: 运行 π θ old \pi_{\theta_{\text{old}}} π θ old ,采集一批轨迹 D \mathcal{D} D
4: 为 D \mathcal{D} D 的每个时间步计算 A ^ t \widehat{A}_t A t
5: for k ← 1 k \gets 1 k ← 1 to K K K do
6: for all minibatch B ⊆ D B \subseteq \mathcal{D} B ⊆ D do
7: θ ← θ + α ∇ θ L C L I P ( θ ; B ) \theta \gets \theta + \alpha \nabla_\theta L^{\mathrm{CLIP}}(\theta; B) θ ← θ + α ∇ θ L CLIP ( θ ; B )
8: end for
9: end for
10: until 收敛
典型取值是 K K K 在 3 3 3 到 10 10 10 之间,ϵ = 0.2 \epsilon = 0.2 ϵ = 0.2 。
第 1 步是本章全部论证兑现的地方。冻住 θ old \theta_{\text{old}} θ old ,就固定了”这批数据被认为来自哪个分布”,而它在整整 K K K 个 epoch 里保持冻结,同时 θ \theta θ 正从它身边走开。内层循环刚开始时 θ = θ old \theta = \theta_{\text{old}} θ = θ old ,于是处处 r t = 1 r_t = 1 r t = 1 ,第一次更新恰好就是上一章那个策略梯度步。此后每一次更新都是重用,用比值付账,由 clip 看管。循环结束时刷新 θ old \theta_{\text{old}} θ old ,这批数据被丢弃。
收益就是那个因子 K K K 。一轮生成现在支撑 K K K 个 epoch 的更新而不是一次,而在一个生成主导成本的场景里,这接近于挂钟时间上的 K K K 倍。
clip 到底保证了什么
到这里为止装配出来的故事很整齐:重用需要局部性,局部性由 K L \mathrm{KL} KL 度量,而 clip 把比值保持在一附近的一个带里。最后这一步里究竟有多少是真的,值得讲精确,因为那个整齐的版本是假的,而且流传甚广。
clip 目标并没有把 r t r_t r t 约束在 [ 1 − ϵ , 1 + ϵ ] [1-\epsilon, 1+\epsilon] [ 1 − ϵ , 1 + ϵ ] 里。它做不到:L C L I P L^{\mathrm{CLIP}} L CLIP 是一个目标函数,不是一个约束集,而上面那条命题只说了它关于该样本自己那个比值 的导数在带外消失。
参数是共享的。真正被施加的梯度是对整个 minibatch 求和,
∇ θ L C L I P = E t [ ∂ L t ∂ r t ∇ θ r t ( θ ) ] , \nabla_\theta L^{\mathrm{CLIP}}
=
\mathbb{E}_t\left[\frac{\partial L_t}{\partial r_t}\,\nabla_\theta r_t(\theta)\right], ∇ θ L CLIP = E t [ ∂ r t ∂ L t ∇ θ r t ( θ ) ] , 而一个自身 ∂ L t / ∂ r t \partial L_t/\partial r_t ∂ L t / ∂ r t 已被掩成零的样本,它的 r t r_t r t 照样在动,因为 minibatch 里其他每一个样本都在拉扯同一个 θ \theta θ 。把一个样本的梯度掩掉,阻止的是目标函数继续推 它出去;这对它被邻居捎带 出去毫无作用。而且任何地方都没有东西会把一个跑偏的比值投影 回来:一旦 r t r_t r t 到了带外,clip 唯一的反应就是不再关心。
于是在 K K K 个 epoch 的 minibatch 更新之后,比值漂出带外是家常便饭,而 K L ‾ ( θ old , θ ) \overline{\mathrm{KL}}(\theta_{\text{old}}, \theta) KL ( θ old , θ ) 不被 ϵ \epsilon ϵ 的任何函数所约束。
这对”本章该怎么读”的后果值得不加掩饰地说出来。PPO 的信赖域读法,即 ϵ \epsilon ϵ 划出一个把策略关在里面的区域,是一个关于目标函数形状的故事,不是一个关于算法行为的定理。真的那部分更弱,但仍然有用:clip 移除了大幅策略改变的回报 ,于是优化器没有动机去寻求它们,而实践中比值大多数时候待在一附近,大多数时候。“没有动机去做”和”做不到”之间的这道缝,正是为什么各种实现都在训练时监控 K L ‾ \overline{\mathrm{KL}} KL 、并在它涨得太大时提前中止内层循环。假如 clip 真做到了人们通常描述的那件事,这道保险就是多余的。
拓展:惩罚变体
提出 PPO 的那篇论文给了两种把策略留在 π θ old \pi_{\theta_{\text{old}}} π θ old 附近的办法,本章只展开了其中一种。另一种把 clip 换成一个显式的惩罚项:
L K L P E N ( θ ) = E t [ r t ( θ ) A ^ t ] − β K L ‾ ( θ old , θ ) , L^{\mathrm{KLPEN}}(\theta)
=
\mathbb{E}_t\big[r_t(\theta)\widehat A_t\big]
-
\beta\, \overline{\mathrm{KL}}(\theta_{\text{old}}, \theta), L KLPEN ( θ ) = E t [ r t ( θ ) A t ] − β KL ( θ old , θ ) ,
其中 β \beta β 在迭代之间自适应:测得的 K L ‾ \overline{\mathrm{KL}} KL 超过目标就调大,低于目标就调小。它是这个想法更直接的表达,因为它惩罚的恰好就是前面几节识别出来的那个量,而且它不假装”约束住一个比值就等于约束住一个散度”。
但它不是这个领域最后选定的那个。clip 变体在原论文的对比里表现更好,也更好调(一个常数,而不是一个带目标值和调整规则的控制器),而后续工作说 PPO 时基本上都是指 clip 目标。惩罚变体在这里作为一个指针而非一个话题:想要它的读者可以在原论文和 Easy RL 第 5.2.1 节 里找到。
延伸阅读
本章假设了策略梯度恒等式、log-derivative 技巧、轨迹分解,以及 advantage 估计量,它们都来自策略梯度 ;还假设了回报、价值与 Bellman 方程 里的价值函数。
这里引入的 K L \mathrm{KL} KL 散度,会在讲偏好类方法的那几章里作为一个承重对象、而不是一个诊断指标再次出现:那时候,一个针对固定参考策略(而不是针对上一次迭代)的惩罚项,会成为被优化的目标本身的一部分。
参考资料
评论