【RL第四篇】广义优势估计-Generalized Advantage Estimation(GAE)
一、前言
这次来盘一盘GAE, 全名Generalized Advantage Estimation,看这个大概能猜到和Advantage估计有关系,为什么要学GAE呢,看下面这张图

因为PPO中会利用GAE来估计Advantage, PPO后面会讲。
二、Reinforce方差大
我们回顾一下 Reinforce的问题
计算策略梯度方差很大:
- 每一步都需要与整体的回报相乘,引入了过去的reward
- 利用蒙特卡洛采样来估计,每个轨迹回报可能被差距会很大,方差也很大
也通过引入baseline来减少方差,但也会引入一些偏差。
方差&偏差的定义,这里有写:https://www.zhihu.com/column/c_1492838352466890752
为了平衡偏差和方差: 广义优势估计 (GAE) 是一种旨在在偏差和方差之间找到更好平衡的技术。它通过使用 TD 误差(时间差分误差)和折扣因子(a discount factor)的组合来估计优势函数来实现这一点。
三、GAE
3.1 优势函数Advantage
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)
此时Vπ(s)V_\pi(s)Vπ(s)为state=s时的预期回报,在PPO中这部分是用一个单独的LLM估计的。
Q函数:
对于state、action下的累积回报:
Qπ(s,a)=Eτ∼π[R(τ)∣s0=s,a0=a]Q^\pi(s, a) = \underset{\tau \sim \pi}{\mathrm{E}}\left[ R(\tau) \mid s_0 = s, a_0 = a \right]Qπ(s,a)=τ∼πE[R(τ)∣s0=s,a0=a]
对于轨迹的回报函数有以下等式:
R(τ)=∑t=0∞γtrtR(\tau) = \sum_{t=0}^{\infty} \gamma^t r_tR(τ)=t=0∑∞γtrt
时间步为♾️,且引入折扣因子决定需要考虑多大程度的未来奖励,折旧因子介于 0 到 1 之间。在极端情况下,γ = 0 表示智能体只关心当前奖励,而 γ = 1 表示所有未来奖励都被考虑在内。折现因子越低,未来回报的价值就越低(考虑的越少)。
如果用价值函数(V)对Q函数进行估计,可以表示为:
Q^π(st,at)=r(st,at)+γVπ(st+1)\hat{Q}^\pi(s_t, a_t) = r(s_t, a_t) + \gamma V^\pi(s_{t+1})Q^π(st,at)=r(st,at)+γVπ(st+1)
此时的即时奖励r(st,at)r(s_t, a_t)r(st,at)与下一个状态的预期回报相加,这里同样有折旧因子。
所以Advantage可以表示为:
A^π(st,at)=[r(st,at)+γVπ(st+1)]−Vπ(st)\hat{A}^\pi(s_t, a_t) = \left[ r(s_t, a_t) + \gamma V^\pi(s_{t+1}) \right] - V^\pi(s_t)A^π(st,at)=[r(st,at)+γVπ(st+1)]−Vπ(st)
3.2 扩展估计
如果把价值函数对Q函数的估计放到到未来k个时间步骤,则有:
Q^π(st,at)=r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)\hat{Q}^\pi(s_t, a_t) = r(s_t, a_t) + \gamma r(s_{t+1}, a_{t+1}) + \ldots + \gamma^{(k-1)} r(s_{t+k-1}, a_{t+k-1}) + \gamma^k V^\pi(s_{t+k})Q^π(st,at)=r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)
则Advantage函数为:
A^π(st,at)=[r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)]−Vπ(st)\hat{A}^\pi(s_t, a_t) = \left[r(s_t, a_t) + \gamma r(s_{t+1}, a_{t+1}) + \ldots + \gamma^{(k-1)} r(s_{t+k-1}, a_{t+k-1}) + \gamma^k V^\pi(s_{t+k}) \right] - V^\pi(s_t)A^π(st,at)=[r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)]−Vπ(st)
对于Advantage的估计,针对k时间步:
我们在估计中包含的时间步长数量会影响这种权衡:
-
更少的时间步长: 更小的方差( 更少的噪声 )但潜在的偏差更高( 更多地依赖于不准确的价值函数V )。
-
更多时间步长: 更多方差( 更多噪音 )但潜在更低偏差( 更多依赖实际奖励 )。
3.3 引入残差
对于VπV^\piVπ也存在等式:
Vπ(st)=r(st,at)+γVπ(st+1)V^\pi(s_t) = r(s_t, a_t) + \gamma V^\pi(s_{t+1})Vπ(st)=r(st,at)+γVπ(st+1)
因为VπV^\piVπ本身是LLM估计的,本身就存在一定的残差,则可以得到残差关于价值函数的等式来表示残差:
δt=r(st,at)−(Vπ(st)−γVπ(st+1))\delta_t = r(s_t, a_t) - \left( V^\pi(s_t) - \gamma V^\pi(s_{t+1}) \right)δt=r(st,at)−(Vπ(st)−γVπ(st+1))
你会神奇的发现残差和优势函数此时是一个方程。也可以很好的理解这个事情:衡量了当前动作带来的回报与预期之间的差距,而优势函数恰恰是计算这个和预期之间的差距。
3.4 优势函数与残差
通过以下三个公式:
- k步回报:Q^π(st,at)=r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)\hat{Q}^\pi(s_t, a_t) = r(s_t, a_t) + \gamma r(s_{t+1}, a_{t+1}) + \ldots + \gamma^{(k-1)} r(s_{t+k-1}, a_{t+k-1}) + \gamma^k V^\pi(s_{t+k})Q^π(st,at)=r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)
- 残差(TD 误差):δt=r(st,at)−(Vπ(st)−γVπ(st+1))\delta_t = r(s_t, a_t) - \left( V^\pi(s_t) - \gamma V^\pi(s_{t+1}) \right)δt=r(st,at)−(Vπ(st)−γVπ(st+1))
- 优势函数:A^π(st,at)=[r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)]−Vπ(st)\hat{A}^\pi(s_t, a_t) = \left[r(s_t, a_t) + \gamma r(s_{t+1}, a_{t+1}) + \ldots + \gamma^{(k-1)} r(s_{t+k-1}, a_{t+k-1}) + \gamma^k V^\pi(s_{t+k}) \right] - V^\pi(s_t)A^π(st,at)=[r(st,at)+γr(st+1,at+1)+…+γ(k−1)r(st+k−1,at+k−1)+γkVπ(st+k)]−Vπ(st)
用残差替换 rtr_trt, rt+1r_{t+1}rt+1等等
利用残差的定义可以得到:
r(st+l,at+l)=δt+l+Vπ(st+l)−γVπ(st+l+1) r(s_{t+l}, a_{t+l})= \delta_{t+l} + V^\pi(s_{t+l}) - \gamma V^\pi(s_{t+l+1}) r(st+l,at+l)=δt+l+Vπ(st+l)−γVπ(st+l+1)
将每一项 rrr 替换:
A^tπ=[δt+V(st)−γV(st+1)]+γ[δt+1+V(st+1)−γV(st+2)]+⋯+γk−1[δt+k−1+V(st+k−1)−γV(st+k)]+γkV(st+k)−V(st) \begin{align*} \hat{A}_t^\pi &= \left[ \delta_t + V(s_t) - \gamma V(s_{t+1}) \right] + \gamma \left[ \delta_{t+1} + V(s_{t+1}) - \gamma V(s_{t+2}) \right] + \cdots \\ &\quad + \gamma^{k-1} \left[ \delta_{t+k-1} + V(s_{t+k-1}) - \gamma V(s_{t+k}) \right] + \gamma^k V(s_{t+k}) - V(s_t) \end{align*} A^tπ=[δt+V(st)−γV(st+1)]+γ[δt+1+V(st+1)−γV(st+2)]+⋯+γk−1[δt+k−1+V(st+k−1)−γV(st+k)]+γkV(st+k)−V(st)
展开后,中间的 V(st+1),V(st+2),…,V(st+k−1)V(s_{t+1}), V(s_{t+2}), \dots, V(s_{t+k-1})V(st+1),V(st+2),…,V(st+k−1) 会被相邻项的正负项抵消:
A^tπ=δt+V(st)−γV(st+1)+γδt+1+γV(st+1)−γ2V(st+2)+⋯+γk−1δt+k−1+γk−1V(st+k−1)−γkV(st+k)+γkV(st+k)−V(st) \begin{align*} \hat{A}_t^\pi &= \delta_t + V(s_t) - \gamma V(s_{t+1}) + \gamma \delta_{t+1} + \gamma V(s_{t+1}) - \gamma^2 V(s_{t+2}) + \cdots \\ &\quad + \gamma^{k-1} \delta_{t+k-1} + \gamma^{k-1} V(s_{t+k-1}) - \gamma^k V(s_{t+k}) + \gamma^k V(s_{t+k}) - V(s_t) \end{align*} A^tπ=δt+V(st)−γV(st+1)+γδt+1+γV(st+1)−γ2V(st+2)+⋯+γk−1δt+k−1+γk−1V(st+k−1)−γkV(st+k)+γkV(st+k)−V(st)
观察发现:
- V(st)V(s_t)V(st) 和 −V(st)-V(s_t)−V(st) 抵消。
- −γV(st+1)-\gamma V(s_{t+1})−γV(st+1) 和 +γV(st+1)+\gamma V(s_{t+1})+γV(st+1) 抵消。
- −γ2V(st+2)-\gamma^2 V(s_{t+2})−γ2V(st+2) 和 +γ2V(st+2)+\gamma^2 V(s_{t+2})+γ2V(st+2) 抵消(以此类推)。
- 最后一项 −γkV(st+k)-\gamma^k V(s_{t+k})−γkV(st+k) 和 +γkV(st+k)+\gamma^k V(s_{t+k})+γkV(st+k) 抵消。
消去所有中间项后,剩余的项为:
A^tπ=δt+γδt+1+γ2δt+2+⋯+γk−1δt+k−1 \hat{A}_t^\pi = \delta_t + \gamma \delta_{t+1} + \gamma^2 \delta_{t+2} + \cdots + \gamma^{k-1} \delta_{t+k-1} A^tπ=δt+γδt+1+γ2δt+2+⋯+γk−1δt+k−1
即:
A^tπ=∑l=1kγl−1δt+l−1=∑l=0kγlδt+l(变量替换 l→l−1) \hat{A}_t^\pi = \sum_{l=1}^k \gamma^{l-1} \delta_{t+l-1} = \sum_{l=0}^k \gamma^l \delta_{t+l} \quad (\text{变量替换 } l \to l-1) A^tπ=l=1∑kγl−1δt+l−1=l=0∑kγlδt+l(变量替换 l→l−1)
3.5 平衡偏差和方差
k选取决定了偏差和方差的关系,如果k太小,则偏差会放大,k太大,方差会放大。
为了平衡优势估计中的偏差 - 方差权衡,广义优势估计(GAE)将优势函数定义为 k 步优势的指数移动平均:
A^tGAE(γ,λ)=(1−λ)(A^t(1)+λA^t(2)+λ2A^t(3)+⋯ )=(1−λ)(δt+λ(δt+γδt+1)+λ2(δt+γδt+1+γ2δt+2)+… )=(1−λ)(δt(1+λ+λ2+… )+γδt+1(λ+λ2+λ3+… )+γ2δt+2(λ2+λ3+λ4+… )+… )=(1−λ)(δt(11−λ)+γδt+1(λ1−λ)+γ2δt+2(λ21−λ)+… )=∑k=0∞(γλ)kδt+k.(9) \begin{align*} \hat{A}_t^{\text{GAE}(\gamma, \lambda)} &= (1 - \lambda)\left( \hat{A}_t^{(1)} + \lambda \hat{A}_t^{(2)} + \lambda^2 \hat{A}_t^{(3)} + \cdots \right) \\ &= (1 - \lambda)\left( \delta_t + \lambda(\delta_t + \gamma \delta_{t+1}) + \lambda^2(\delta_t + \gamma \delta_{t+1} + \gamma^2 \delta_{t+2}) + \dots \right) \\ &= (1 - \lambda)\left( \delta_t(1 + \lambda + \lambda^2 + \dots) + \gamma \delta_{t+1}(\lambda + \lambda^2 + \lambda^3 + \dots) \right. \\ &\quad \left. + \gamma^2 \delta_{t+2}(\lambda^2 + \lambda^3 + \lambda^4 + \dots) + \dots \right) \\ &= (1 - \lambda)\left( \delta_t \left( \frac{1}{1 - \lambda} \right) + \gamma \delta_{t+1} \left( \frac{\lambda}{1 - \lambda} \right) + \gamma^2 \delta_{t+2} \left( \frac{\lambda^2}{1 - \lambda} \right) + \dots \right) \\ &= \sum_{k=0}^{\infty} (\gamma \lambda)^k \delta_{t+k}. \end{align*} \tag{9} A^tGAE(γ,λ)=(1−λ)(A^t(1)+λA^t(2)+λ2A^t(3)+⋯)=(1−λ)(δt+λ(δt+γδt+1)+λ2(δt+γδt+1+γ2δt+2)+…)=(1−λ)(δt(1+λ+λ2+…)+γδt+1(λ+λ2+λ3+…)+γ2δt+2(λ2+λ3+λ4+…)+…)=(1−λ)(δt(1−λ1)+γδt+1(1−λλ)+γ2δt+2(1−λλ2)+…)=k=0∑∞(γλ)kδt+k.(9)
这个公式的推导有两个关键:
等比数列求和等式
对于首项为 aaa、公比为 rrr 的等比数列,前kkk 项和为:
Sk=a+ar+ar2+⋯+ark−1=a⋅1−rk1−r(r≠1) S_k = a + ar + ar^2 + \cdots + ar^{k-1} = a \cdot \frac{1 - r^k}{1 - r} \quad (r \neq 1) Sk=a+ar+ar2+⋯+ark−1=a⋅1−r1−rk(r=1)
利用等比数列有
1+λ+λ2+⋯+λk−1=1−λk1−λ 1 + \lambda + \lambda^2 + \cdots + \lambda^{k-1} = \frac{1 - \lambda^k}{1 - \lambda} 1+λ+λ2+⋯+λk−1=1−λ1−λk
其中λ\lambdaλ为 0-1的数字,k趋近于无穷,所以λk\lambda^kλk趋近于0
3.6 参数λ\lambdaλ对方差和偏差的影响
λ\lambdaλ是调节偏差和方差的重要超参数
- 值越大,则考虑后续步骤更多,偏差更小,但方差更大;
- 值越小,则考虑后续步骤更少,更多依赖当前价值预估,偏差更大,但方差更小。
当为0的时候,则偏差最大:
GAE(γ,0):A^t=δt=rt+γV(st+1)−V(st). \text{GAE}(\gamma, 0) : \hat{A}_t = \delta_t = r_t + \gamma V(s_{t+1}) - V(s_t). GAE(γ,0):A^t=δt=rt+γV(st+1)−V(st).
当为1的时候,则方差最大:
GAE(γ,1):A^t=∑k=0∞γkδt+1=∑k=0∞γkrt+1−V(st). \text{GAE}(\gamma, 1) : \hat{A}_t = \sum_{k=0}^{\infty} \gamma^k \delta_{t+1} = \sum_{k=0}^{\infty} \gamma^k r_{t+1} - V(s_t). GAE(γ,1):A^t=k=0∑∞γkδt+1=k=0∑∞γkrt+1−V(st).
Ref
- https://arxiv.org/pdf/2501.03262
- https://nn.labml.ai/zh/rl/ppo/gae.html
- https://shivang-ahd.medium.com/generalized-advantage-estimation-a-deep-dive-into-bias-variance-and-policy-gradients-a5e0b3454dad
- https://arxiv.org/pdf/2307.04964
更多推荐

所有评论(0)