Direct policy differentiation
用带参数 \(\theta\) 的策略 \(\pi_\theta\) 与环境交互,会得到一条轨迹 \(\tau\)。记这套策略产生轨迹 \(\tau\) 的概率或密度为 \(\pi_\theta(\tau)\),整条轨迹的总奖励为 \(r(\tau)=\sum_{t=1}^T r(s_t,a_t)\)。我们想最大化的是这套策略反复执行时的平均总奖励:
\[
\mathcal{J}(\theta)
=\mathbb{E}_{\tau\sim\pi_\theta}[r(\tau)]
=\int\pi_\theta(\tau)r(\tau)\,\mathrm{d}\tau.
\]
\(\mathcal{J}(\theta)\) 衡量当前策略有多好;为了更新参数,我们还需要求它的梯度 \(\nabla_\theta\mathcal{J}(\theta)\)。这里假设环境的转移和奖励规则不依赖 \(\theta\):固定一条轨迹后,\(r(\tau)\) 就固定了,但改变策略参数会改变各条轨迹出现的概率。因此,求导时需要处理的是期望中的分布 \(\pi_\theta\)。
在积分与求导可以交换的条件下,REINFORCE Algorithm (Williams, 1992) 使用下面的梯度表达式:
\[
\begin{aligned}
\nabla_\theta \mathcal{J}(\theta)
&= \nabla_\theta\mathbb{E}_{\tau\sim\pi_\theta}[r(\tau)] \\
&= \int \nabla_\theta \pi_\theta(\tau) r(\tau)\,\mathrm{d}\tau \\
&= \int \pi_\theta(\tau)\frac{\nabla_\theta\pi_\theta(\tau)}{\pi_\theta(\tau)}r(\tau)\,\mathrm{d}\tau \\
&= \int \pi_\theta(\tau)r(\tau)\nabla_\theta\log\pi_\theta(\tau)\,\mathrm{d}\tau \\
&= \mathbb{E}_{\tau\sim\pi_\theta}\left[r(\tau)\cdot\nabla_\theta\log\pi_\theta(\tau)\right] \\
&= \mathbb{E}_{\tau\sim\pi_\theta}\left[
\left(\sum_{t=1}^T\nabla_\theta\log\pi_\theta(a_t\mid s_t)\right)
\left(\sum_{t=1}^T r(s_t,a_t)\right)
\right]
\end{aligned}
\]
开头的期望等于 \(\mathcal{J}(\theta)\),推导末尾的期望等于 \(\nabla_\theta\mathcal{J}(\theta)\)。 两者都按同一个 \(\pi_\theta\) 取平均,但被平均的量不同:前者是奖励 \(r(\tau)\),后者是梯度贡献 \(r(\tau)\nabla_\theta\log\pi_\theta(\tau)\)。前者给出策略的表现,后者给出这个表现随参数变化的方向和幅度。
把梯度重新写成期望,是为了方便采样估计。对同一批轨迹 \(\tau_1,\ldots,\tau_N\sim\pi_\theta\),分别取下面两个平均,就能估计目标值和目标梯度:
\[
\begin{aligned}
\hat{\mathcal{J}}&=\frac1N\sum_{i=1}^N r(\tau_i),\\
\hat g&=\frac1N\sum_{i=1}^N r(\tau_i)\nabla_\theta\log\pi_\theta(\tau_i).
\end{aligned}
\]
所以说,在 \(\tau\sim\pi_\theta\) 时,\(r(\tau)\cdot\nabla_\theta\log\pi_\theta(\tau)\) 是目标梯度 \(\nabla_\theta\mathcal{J}(\theta)\) 的无偏估计。这里需要的是 reward 加权的 log probability 的梯度,单独的 \(\log\pi_\theta(\tau)\) 并不是这个估计量。可惜这玩意儿的 variance 很高。
从轨迹积分到代码里的 loss
按 \(\pi_\theta\) 采样后,直接平均 \(r(\tau)\) 就能估计 \(\mathcal{J}(\theta)\)。积分里的概率权重已经由采样频率体现了;如果再平均 \(\pi_\theta(\tau)r(\tau)\),就会多乘一次轨迹概率。
但采样后的 reward 是固定数据,直接对它求导无法反映轨迹分布随参数的变化。为了用梯度下降实现正文推导出的更新,构造下面的 loss,并在求导时固定采到的轨迹和奖励:
\[
\begin{aligned}
\ell_\theta(\tau)&=-r(\tau)\log\pi_\theta(\tau),\\
-\nabla_\theta\ell_\theta(\tau)
&=r(\tau)\nabla_\theta\log\pi_\theta(\tau).
\end{aligned}
\]
因此,loss 的负梯度 \(-\nabla_\theta\ell_\theta(\tau)\) 是 policy gradient 的无偏估计,loss 的数值本身并不估计负的期望奖励。Spinning Up (Achiam, 2018) 强调,这个无偏保证是在采样所用的策略参数处成立的;更新参数后,一直复用旧 batch 最小化它,并没有同样的保证。
对一条轨迹,代码中把各步实际采到的 action 的 log probability 相加:
trajectory_log_prob = action_log_probs.sum()
loss = -trajectory_return.detach() * trajectory_log_prob
这里 trajectory_return 是总奖励,求导时作为常量;action_log_probs 保留策略参数的计算图。对一个 batch,再平均各条轨迹的 loss。完整轨迹概率中的初始状态和环境转移项在本文假设下与 \(\theta\) 无关,因此计算梯度时可以省略。
Reduce Variance
仔细观察这个式子,其实这玩意儿是 MLE 那个梯度对 \(r(s_t,a_t)\) 加权了:
\[
\nabla_\theta\mathcal{J}(\theta)
= \mathbb{E}_{\tau\sim\pi_\theta}\left[
\sum_{t=0}^T\Psi_t\nabla_\theta\log\pi_\theta(a_t\mid s_t)
\right]
\]
当前我们是有:\(\Psi_t=\sum_{t'=0}^T r(s_{t'},a_{t'})\)
Don't Let the Past Distract You
一种简单的方法来减小 variance,是我们令 \(\Psi_t=\sum_{t'=t}^T r(s_{t'},a_{t'})\)
因为其实对于 \(a_t\) 来说,他做啥对于 \(t\) 之前的 reward 来说是不具有参考价值的。因此我们主要考虑后面的 reward。这玩意儿直觉上挺清楚的,但数学上想了半天才想明白为啥是对的。主要参考了这篇文章 (Achiam, 2018)。
证明的不造为啥让我想起了 MLE。主要用到的就是一个叫做 EGLP lemma 的东西(其实好像用这个 lemma 需要积分和导数的可交换性,貌似 (Hogg et al., 2013) 里写挺详细的):
\[
\begin{aligned}
\mathbb{E}_{x\sim\mathbb{P}_\theta}\left[\nabla_\theta\log\mathbb{P}_\theta(x)\right]
&= \int\mathbb{P}_\theta(x)\nabla_\theta\log\mathbb{P}_\theta(x)\,\mathrm{d}x \\
&= \int\mathbb{P}_\theta(x)\frac{\nabla_\theta\mathbb{P}_\theta(x)}{\mathbb{P}_\theta(x)}\,\mathrm{d}x \\
&= \nabla_\theta\int\mathbb{P}_\theta(x)\,\mathrm{d}x=0
\end{aligned}
\]
其实跟 MLE 是一样的嘛:
\[
\mathbb{E}\left[\frac{\partial}{\partial\theta}\log\mathcal{L}(x\mid\theta)\right]=0
\]
其实我们是要证明嘟是:
\[
\mathbb{E}_{\tau\sim\pi_\theta}\left[
\sum_{t=0}^T\sum_{t'<t}r(s_{t'},a_{t'})\nabla_\theta\log\pi_\theta(a_t\mid s_t)
\right]=0
\]
也就是要证明当 \(t'<t\) 这个时候:
\[
\mathbb{E}_{s_t,a_t,s_{t'},a_{t'}\sim\pi_\theta}\left[
r(s_{t'},a_{t'})\nabla_\theta\log\pi_\theta(a_t\mid s_t)
\right]=0
\]
那么中心思想其实就是咋来区分 \(t'<t\) 捏,我们考虑 \(t'<t\) 是先 reward,再选择:
\[
\mathbb{E}_{s_{t'},a_{t'}\sim\pi_\theta}\left[
r(s_{t'},a_{t'})\cdot
\mathbb{E}_{s_t,a_t\sim\pi_\theta(\cdot\mid s_{t'},a_{t'})}\left[
\nabla_\theta\log\pi_\theta(a_t\mid s_t)\mid s_{t'},a_{t'}
\right]
\right]
\]
其实也就是当 \(r(s_{t'},a_{t'})\) 不依赖于 \(s_t,a_t\) 的时候,本身这个 \(\mathbb{E}[\nabla_\theta\log\pi_\theta(a_t\mid s_t)]\) 他就是 \(0\)。
所以说最终结果是整个期望 \(0\)。
Introducing Baselines
另一个优化是我们考虑加入 baseline。这个直觉就更对了。就是我们考虑把 \(r(s,a)\) 替换成 \(r(s,a)-b\)。因为 reward 这种东西,大家一起加多少减多少肯定都是无所谓的。
当然从数学上来讲也是 EGLP 用用易证的,这里就不多写了。但是我们 baseline 设多少最好呢?从直觉上来讲,是这个 \(b\) 让整个 \(r\) 尽量居中。接下来我们从数学上进行考虑。
我们考虑
\[
\nabla_\theta\mathcal{J}(\theta)
= \mathbb{E}_{\tau\sim\pi_\theta}\left[
\nabla_\theta\log\pi_\theta(\tau)\cdot(r(\tau)-b)
\right]
\]
的方差(对于向量梯度,这里取各分量方差之和)
\[
\begin{aligned}
\sigma^2
&= \mathbb{E}_{\tau\sim\pi_\theta}\left[
\left\|\nabla_\theta\log\pi_\theta(\tau)\cdot(r(\tau)-b)\right\|^2
\right] \\
&\quad -\left\|\mathbb{E}_{\tau\sim\pi_\theta}\left[
\nabla_\theta\log\pi_\theta(\tau)\cdot(r(\tau)-b)
\right]\right\|^2 \\
&= \mathbb{E}_{\tau\sim\pi_\theta}\left[
\left\|\nabla_\theta\log\pi_\theta(\tau)\cdot(r(\tau)-b)\right\|^2
\right] \\
&\quad -\left\|\mathbb{E}_{\tau\sim\pi_\theta}\left[
\nabla_\theta\log\pi_\theta(\tau)\cdot r(\tau)
\right]\right\|^2
\end{aligned}
\]
我们解
\[
\frac{\partial}{\partial b}\sigma^2
= \frac{\partial}{\partial b}\mathbb{E}_{\tau\sim\pi_\theta}\left[
\left\|\nabla_\theta\log\pi_\theta(\tau)\cdot(r(\tau)-b)\right\|^2
\right]=0
\]
可以得到
\[
b=\frac{
\mathbb{E}_{\tau\sim\pi_\theta}\left[\left\|\nabla_\theta\log\pi_\theta(\tau)\right\|^2\cdot r(\tau)\right]
}{
\mathbb{E}_{\tau\sim\pi_\theta}\left[\left\|\nabla_\theta\log\pi_\theta(\tau)\right\|^2\right]
}
\]
这啥捏,这其实是 reward 的加权期望。
但其实这个 baseline 挺难算的,所以我们通常不会用这个最优的 baseline。而是去找一个相对比较好的。
Off-Policy Policy Gradients
之前我们做的都是 on-policy 的。但真实在训练的时候,我们很难做到稍稍改一点 \(\theta\),就重新生成一堆新的 \(\tau\)。这样是非常 inefficient 的。
所以现在可能的问题是,我们没有关于 \(\tau\sim\pi_\theta\) 的数据,但是我们可能有一个其他的 distribution,通过这个 distribution 来 sample 出的数据。也就是 \(\tau\sim\overline{\pi}\)。
我们需要使用的一个 trick 叫做 importance sampling:
\[
\begin{aligned}
\mathbb{E}_{x\sim p(x)}[f(x)]
&= \int p(x)f(x)\,\mathrm{d}x \\
&= \int q(x)\frac{p(x)}{q(x)}f(x)\,\mathrm{d}x \\
&= \mathbb{E}_{x\sim q(x)}\left[\frac{p(x)}{q(x)}f(x)\right]
\end{aligned}
\]
所以说我们的 RL objective 可以改成:
\[
\begin{aligned}
\mathcal{J}(\theta)
&= \mathbb{E}_{\tau\sim\overline{\pi}}\left[
\frac{\pi_\theta(\tau)}{\overline{\pi}(\tau)}r(\tau)
\right] \\
&= \mathbb{E}_{\tau\sim\overline{\pi}}\left[
\frac{p(s_1)\prod_{t=1}^T\pi_\theta(a_t\mid s_t)p(s_{t+1}\mid s_t,a_t)}
{p(s_1)\prod_{t=1}^T\overline{\pi}(a_t\mid s_t)p(s_{t+1}\mid s_t,a_t)}
r(\tau)
\right] \\
&= \mathbb{E}_{\tau\sim\overline{\pi}}\left[
r(\tau)\prod_{t=1}^T\frac{\pi_\theta(a_t\mid s_t)}{\overline{\pi}(a_t\mid s_t)}
\right]
\end{aligned}
\]
所以说我们可以推导梯度
\[
\begin{aligned}
\nabla_\theta\mathcal{J}(\theta)
&= \mathbb{E}_{\tau\sim\overline{\pi}}\left[
\frac{\pi_\theta(\tau)}{\overline{\pi}(\tau)}\nabla_\theta\log\pi_\theta(\tau)r(\tau)
\right] \\
&= \mathbb{E}_{\tau\sim\overline{\pi}}\left[
\left(\prod_{t=1}^T\frac{\pi_\theta(a_t\mid s_t)}{\overline{\pi}(a_t\mid s_t)}\right)
\left(\sum_{t=1}^T\nabla_\theta\log\pi_\theta(a_t\mid s_t)\right)
\left(\sum_{t=1}^T r(s_t,a_t)\right)
\right] \\
&= \mathbb{E}_{\tau\sim\overline{\pi}}\Biggl[
\sum_{t=1}^T\nabla_\theta\log\pi_\theta(a_t\mid s_t)
\left(\prod_{t'=1}^t\frac{\pi_\theta(a_{t'}\mid s_{t'})}{\overline{\pi}(a_{t'}\mid s_{t'})}\right) \\
&\qquad\qquad\cdot\left(
\sum_{t'=t}^T r(s_{t'},a_{t'})
\prod_{t''=t+1}^{t'}\frac{\pi_\theta(a_{t''}\mid s_{t''})}{\overline{\pi}(a_{t''}\mid s_{t''})}
\right)
\Biggr]
\end{aligned}
\]
References
Achiam, J. (2018).
Spinning Up in Deep Reinforcement Learning.
spinningup.openai.com
Hogg, R. V., McKean, J. W., & Craig, A. T. (2013).
Introduction to Mathematical Statistics (7th ed.). Pearson.
scholarworks.wmich.edu
Williams, R. J. (1992). Simple statistical gradient-following algorithms for connectionist reinforcement learning.
Machine Learning,
8(3–4), 229–256.
doi.org