一、前言

这次来盘一盘GAE, 全名Generalized Advantage Estimation,看这个大概能猜到和Advantage估计有关系,为什么要学GAE呢,看下面这张图

因为PPO中会利用GAE来估计Advantage, PPO后面会讲。

二、Reinforce方差大

我们回顾一下 Reinforce的问题

计算策略梯度方差很大:

  1. 每一步都需要与整体的回报相乘,引入了过去的reward
  2. 利用蒙特卡洛采样来估计,每个轨迹回报可能被差距会很大,方差也很大

也通过引入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)++γ(k1)r(st+k1,at+k1)+γ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)++γ(k1)r(st+k1,at+k1)+γ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)++γ(k1)r(st+k1,at+k1)+γ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)++γ(k1)r(st+k1,at+k1)+γ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)]++γk1[δt+k1+V(st+k1)γ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+k1) 会被相邻项的正负项抵消:

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)++γk1δt+k1+γk1V(st+k1)γ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++γk1δt+k1

即:

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=1kγl1δt+l1=l=0kγlδt+l(变量替换 ll1)

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++ark1=a1r1rk(r=1)

利用等比数列有

1+λ+λ2+⋯+λk−1=1−λk1−λ 1 + \lambda + \lambda^2 + \cdots + \lambda^{k-1} = \frac{1 - \lambda^k}{1 - \lambda} 1+λ+λ2++λk1=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+1V(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
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐