0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

ニュートン法とPPO

0
Posted at

1. ニュートン法(Newton's Method)

1.1 損失の二次近似

損失 $\mathcal L$ を最小化するパラメータ $\theta$ を求める。

\theta\in\mathbb R^P,\qquad g=\nabla_\theta\mathcal L(\theta),\qquad H=\nabla_\theta^2\mathcal L(\theta),\qquad\mathcal L\in C^2
\mathcal L(\theta+d)\approx\mathcal L(\theta)+g^\top d+\frac12d^\top Hd
\begin{aligned}
H\succ0\quad&\Longrightarrow\quad
\Delta\theta=\arg\min_d\left[g^\top d+\frac12d^\top Hd\right]\\
&\Longrightarrow\quad g+H\Delta\theta=0\\
&\Longrightarrow\quad\Delta\theta=-H^{-1}g
\end{aligned}

1.2 正定値性とステップ長

$H$ は正定値とは限らないので、$\lambda I$ を加える。

\begin{aligned}
B&=H+\lambda I\succ0,\qquad \lambda\ge0,\\
d&=-B^{-1}g,\\
g^\top d&=-g^\top B^{-1}g<0\qquad(g\ne0)
\end{aligned}
\Delta\theta=\alpha d=-\alpha(H+\lambda I)^{-1}g,\qquad\alpha>0

2. 自然勾配法:分布の変化で更新を測る

2.1 KLの二次近似とFisher情報行列

条件 $x\sim\rho$ ごとに、$\theta$ が分布 $q_\theta(y\mid x)$ を決める。分布の変化をKLダイバージェンスで測る。

\mathcal K_\theta(d):=\mathbb E_{x\sim\rho}
D_{\mathrm{KL}}\!\left(q_\theta(\cdot\mid x)\middle\|q_{\theta+d}(\cdot\mid x)\right),\qquad
\mathbb E_\theta[\cdot]:=\mathbb E_{x\sim\rho,\;y\sim q_\theta(\cdot\mid x)}[\cdot]

台が局所的に共通で、微分と期待値を交換できるとする。

s_\theta:=\nabla_\theta\log q_\theta(y\mid x),\qquad
\mathbb E_{y\sim q_\theta}[s_\theta]=\nabla_\theta\int q_\theta(y\mid x)\,\mathrm dy=0,\qquad
F:=\mathbb E_\theta[s_\theta s_\theta^\top]\succeq0
\mathcal K_\theta(0)=0,\qquad
\left.\nabla_d\mathcal K_\theta\right|_{d=0}=-\mathbb E_\theta[s_\theta]=0,\qquad
\left.\nabla_d^2\mathcal K_\theta\right|_{d=0}=-\mathbb E_\theta[\nabla_\theta^2\log q_\theta]=F
\mathcal K_\theta(d)=\frac12d^\top Fd+o(\|d\|^2)

2.2 自然勾配の導出

損失の一次改善と、分布変化の二次ペナルティを釣り合わせる。

\begin{aligned}
\Delta\theta
&=\arg\min_d\left[g^\top d+\frac1{2\eta}d^\top(F+\lambda I)d\right],\qquad\eta>0,\ \lambda>0,\\
0&=g+\frac1\eta(F+\lambda I)\Delta\theta,\\
\Delta\theta&=-\eta(F+\lambda I)^{-1}g
\end{aligned}
\begin{aligned}
H&=\nabla_\theta^2\mathcal L(\theta)&&\text{:損失の曲率}\\
F&=\left.\nabla_d^2\mathcal K_\theta(d)\right|_{d=0}&&\text{:分布の変化}
\end{aligned}

3. 強化学習:方策が行動の分布を決める

第1〜2節で、更新を測る行列として、損失の曲率 $H$ と分布の変化 $F$ の二つを得た。強化学習では方策そのものが分布なので、$F$ を使える。

3.1 マルコフ決定過程と方策

状態と行動は有限集合とする。連続の場合は和を積分に置き換える。

\mathcal M=(\mathcal S,\mathcal A,P,r,p_0,\gamma),\qquad
s_0\sim p_0,\qquad s_{t+1}\sim P(\cdot\mid s_t,a_t),\qquad r_t=r(s_t,a_t),\qquad 0\le\gamma<1
a_t\sim\pi_\theta(\cdot\mid s_t),\qquad
p_\theta(\tau)=p_0(s_0)\prod_{t\ge0}\pi_\theta(a_t\mid s_t)P(s_{t+1}\mid s_t,a_t),\qquad
\tau=(s_0,a_0,s_1,a_1,\ldots)

報酬は有界とする。

|r(s,a)|\le r_{\max}\quad\Longrightarrow\quad
J(\theta):=\mathbb E_{\tau\sim p_\theta}\left[\sum_{t\ge0}\gamma^tr_t\right],\qquad|J(\theta)|\le\frac{r_{\max}}{1-\gamma}

報酬の和は最大化するので、第1〜2節の損失を $\mathcal L=-J$ と読む。以降の $g$ は $J$ の勾配とし、更新式の符号を反転する。

g:=\nabla_\theta J(\theta)=-\nabla_\theta\mathcal L,\qquad
\Delta\theta_{\mathrm{Natural}}=\eta(F+\lambda I)^{-1}g

3.2 価値関数とアドバンテージ

\begin{aligned}
V^\theta(s)&=\mathbb E_{\tau\sim p_\theta}\left[\sum_{t\ge0}\gamma^tr_t\,\middle|\,s_0=s\right],\\
Q^\theta(s,a)&=r(s,a)+\gamma\,\mathbb E_{s'\sim P(\cdot\mid s,a)}\left[V^\theta(s')\right],\\
A^\theta(s,a)&=Q^\theta(s,a)-V^\theta(s)
\end{aligned}
V^\theta(s)=\mathbb E_{a\sim\pi_\theta(\cdot\mid s)}\left[Q^\theta(s,a)\right]
\quad\Longrightarrow\quad
\mathbb E_{a\sim\pi_\theta(\cdot\mid s)}\left[A^\theta(s,a)\right]=0

割引した状態の訪問分布 $d^\theta$ を定める。第2節の $x\sim\rho$ と違い、$d^\theta$ は $\theta$ に依存する。

d^\theta(s):=(1-\gamma)\sum_{t\ge0}\gamma^t\Pr\nolimits_{p_\theta}(s_t=s),\qquad\sum_{s}d^\theta(s)=1,\qquad
\mathbb E_\theta[\cdot]:=\mathbb E_{s\sim d^\theta,\;a\sim\pi_\theta(\cdot\mid s)}[\cdot]

3.3 MLPによる方策

MLP $\hat f$ の出力 $z$ を、固定した写像 $\kappa$ で行動の分布にする。

s\xrightarrow{\ \hat f(\cdot;\theta)\ }z\in\mathbb R^m\xrightarrow{\ \kappa\ }\pi_\theta(\cdot\mid s),\qquad
\pi_\theta(a\mid s)=\kappa\!\left(a;\hat f(s;\theta)\right)
\kappa(a;z)=
\begin{cases}
\operatorname{softmax}(z)_a&\text{離散行動},\\
\mathcal N\!\left(a;z,\operatorname{diag}(\sigma^2)\right),\ \ \sigma\ \text{は固定}&\text{連続行動}
\end{cases}

行動1個あたりのスコアとFisherを定める。出力 $z$ についてのFisher $F_z$ は、MLPの記事と同じ形になる。

\psi_\theta(s,a):=\nabla_\theta\log\pi_\theta(a\mid s),\qquad
F_s:=\mathbb E_{a\sim\pi_\theta(\cdot\mid s)}\left[\psi_\theta\psi_\theta^\top\right]
F_z=
\begin{cases}
\operatorname{diag}(q)-qq^\top,\ \ q=\operatorname{softmax}(z)&\text{離散行動},\\
\operatorname{diag}(\sigma^{-2})&\text{連続行動}
\end{cases}

4. 方策勾配

方策と目的 $J$ が決まった。次は、更新の一次の情報 $g$ を方策のスコアで書く。

4.1 軌道の対数尤度

$p_0$ と $P$ は $\theta$ に依存しないので、軌道のスコアは方策のスコアの和になる。

\nabla_\theta\log p_\theta(\tau)=\sum_{t\ge0}\nabla_\theta\log\pi_\theta(a_t\mid s_t)=\sum_{t\ge0}\psi_\theta(s_t,a_t)
\nabla_\theta J
=\int\nabla_\theta p_\theta(\tau)\sum_{t'}\gamma^{t'}r_{t'}\,\mathrm d\tau
=\mathbb E_{\tau\sim p_\theta}\left[\sum_{t\ge0}\psi_\theta(s_t,a_t)\sum_{t'\ge0}\gamma^{t'}r_{t'}\right]

$t'<t$ の報酬は、$s_t$ までの履歴で条件付けると $a_t$ と独立になる。

t'<t\quad\Longrightarrow\quad
\mathbb E\left[\psi_\theta(s_t,a_t)\,r_{t'}\right]
=\mathbb E\left[r_{t'}\,\mathbb E_{a_t\sim\pi_\theta(\cdot\mid s_t)}[\psi_\theta(s_t,a_t)]\right]=0
\mathbb E\left[\sum_{t'\ge t}\gamma^{t'-t}r_{t'}\,\middle|\,s_t,a_t\right]=Q^\theta(s_t,a_t)
\quad\Longrightarrow\quad
\nabla_\theta J=\mathbb E_{\tau\sim p_\theta}\left[\sum_{t\ge0}\gamma^t\,\psi_\theta(s_t,a_t)\,Q^\theta(s_t,a_t)\right]

4.2 方策勾配定理とベースライン

時刻の和を $d^\theta$ にまとめる(Sutton et al., 2000)。

\nabla_\theta J=\frac1{1-\gamma}\mathbb E_\theta\left[Q^\theta(s,a)\,\psi_\theta(s,a)\right]

状態だけの関数 $b(s)$ を引いても期待値は変わらない。

\mathbb E_{a\sim\pi_\theta(\cdot\mid s)}\left[b(s)\,\psi_\theta(s,a)\right]=b(s)\cdot0=0

$b=V^\theta$ と置くと、アドバンテージの形になる。

g=\nabla_\theta J=\frac1{1-\gamma}\mathbb E_\theta\left[A^\theta(s,a)\,\psi_\theta(s,a)\right]

5. 方策の更新を二次で測る

改善する方向 $g$ は分かった。進む幅を決めるには二次の情報が要る。$J$ の二次には状態分布 $d^\theta$ の変化が入るので、まず分布を固定した近似を作る。

5.1 性能差分とサロゲート目的

$\theta'$ の軌道上で $V^\theta$ の差を足し合わせると、和は $V^\theta(s_0)$ だけを残して打ち消し合う(Kakade and Langford, 2002)。

\begin{aligned}
\sum_{t\ge0}\gamma^tr_t
&=V^\theta(s_0)+\sum_{t\ge0}\gamma^t\left(r_t+\gamma V^\theta(s_{t+1})-V^\theta(s_t)\right),\\
\mathbb E\left[r_t+\gamma V^\theta(s_{t+1})-V^\theta(s_t)\,\middle|\,s_t,a_t\right]&=A^\theta(s_t,a_t)
\end{aligned}
J(\theta')-J(\theta)
=\mathbb E_{\tau\sim p_{\theta'}}\left[\sum_{t\ge0}\gamma^tA^\theta(s_t,a_t)\right]
=\frac1{1-\gamma}\mathbb E_{s\sim d^{\theta'},\;a\sim\pi_{\theta'}(\cdot\mid s)}\left[A^\theta(s,a)\right]

$d^{\theta'}$ は $\theta'$ ごとに集め直す必要がある。手元の $d^\theta$ で代用し、行動は重要度重みで $\pi_\theta$ からの標本に直す。softmaxとGaussianは台が全体なので、比は常に定義できる。

\rho_{\theta'}(s,a):=\frac{\pi_{\theta'}(a\mid s)}{\pi_\theta(a\mid s)},\qquad
\mathbb E_{a\sim\pi_{\theta'}}\left[A^\theta(s,a)\right]=\mathbb E_{a\sim\pi_\theta}\left[\rho_{\theta'}(s,a)A^\theta(s,a)\right]
L_\theta(\theta'):=\frac1{1-\gamma}\mathbb E_\theta\left[\rho_{\theta'}(s,a)\,A^\theta(s,a)\right]

$\nabla_{\theta'}\rho_{\theta'}=\rho_{\theta'}\nabla_{\theta'}\log\pi_{\theta'}$ から、$L_\theta$ は $\theta'=\theta$ で $J(\theta')-J(\theta)$ と一次まで一致する。

L_\theta(\theta)=\frac1{1-\gamma}\mathbb E_\theta[A^\theta]=0,\qquad
\left.\nabla_{\theta'}L_\theta(\theta')\right|_{\theta'=\theta}=\frac1{1-\gamma}\mathbb E_\theta\left[A^\theta\psi_\theta\right]=g
J(\theta+d)-J(\theta)=L_\theta(\theta+d)+o(\|d\|)

5.2 サロゲートのHessian

\left.\nabla_{\theta'}^2\rho_{\theta'}\right|_{\theta'=\theta}
=\frac{\nabla_\theta^2\pi_\theta}{\pi_\theta}
=\nabla_\theta^2\log\pi_\theta+\psi_\theta\psi_\theta^\top
H_L:=\left.\nabla_{\theta'}^2L_\theta(\theta')\right|_{\theta'=\theta}
=\frac1{1-\gamma}\mathbb E_\theta\left[A^\theta(s,a)\left(\nabla_\theta^2\log\pi_\theta(a\mid s)+\psi_\theta\psi_\theta^\top\right)\right]

括弧内の行列は、行動について平均すると0になる。$H_L$ は $A^\theta$ と括弧内の行列の共分散であり、$A^\theta$ の符号で正にも負にもなる。

\mathbb E_{a\sim\pi_\theta(\cdot\mid s)}\left[\nabla_\theta^2\log\pi_\theta+\psi_\theta\psi_\theta^\top\right]
=\sum_a\nabla_\theta^2\pi_\theta(a\mid s)=\nabla_\theta^2\,1=0

$L_\theta$ と $J$ の一致は一次までなので、$d^{\theta'}$ の変化が $\nabla^2J$ に入り、一般に $H_L\ne\nabla_\theta^2J$ となる。

2本腕バンディットでの比較

状態1個、$\gamma=0$ とする。$d^\theta$ が $\theta$ に依存しないので、$L_\theta(\theta')=J(\theta')-J(\theta)$ が厳密に成り立ち、$H_L=\nabla^2J$ となる。

\mathcal A=\{1,2\},\qquad x:=\theta_1-\theta_2,\qquad q_1=\pi_\theta(1)=S(x),\qquad q_2=1-q_1,\qquad J=q_1r_1+q_2r_2

MLPの記事のシグモイドの微分 $S'=S(1-S)$、$S''=S(1-S)(1-2S)$ を使う。

\begin{aligned}
J'(x)&=(r_1-r_2)S'(x)=(r_1-r_2)q_1q_2,\\
J''(x)&=(r_1-r_2)S''(x)=(r_1-r_2)q_1q_2(q_2-q_1),\\
\partial_x\log\pi(1)&=q_2,\qquad\partial_x\log\pi(2)=-q_1,\qquad F(x)=q_1q_2^2+q_2q_1^2=q_1q_2
\end{aligned}

良い行動を選ぶ確率が 1/2 未満のとき、$J$ は $x$ について凸になり、ニュートン法は悪い行動のほうへ進む。

r_1>r_2,\ q_1<q_2\quad\Longrightarrow\quad J'>0,\ J''>0\quad\Longrightarrow\quad
\Delta x_{\mathrm{Newton}}=-\frac{J'}{J''}=-\frac1{q_2-q_1}<0

自然勾配は $q_1$ によらず、良い行動のほうへ一定の幅で進む。

\Delta x_{\mathrm{Natural}}=\eta\,\frac{J'}{F}=\eta\,(r_1-r_2)>0

5.3 方策のFisher情報行列

第2節の $x\sim\rho$ を $s\sim d^\theta$ に、$y$ を $a$ に置き換える。$d^\theta$ は現在の $\theta$ で固定する。この $F$ による更新が方策の自然勾配である(Kakade, 2001)。

F:=\mathbb E_\theta\left[\psi_\theta\psi_\theta^\top\right]=\mathbb E_{s\sim d^\theta}[F_s]\succeq0
\bar D(\theta,\theta'):=\mathbb E_{s\sim d^\theta}
D_{\mathrm{KL}}\!\left(\pi_\theta(\cdot\mid s)\middle\|\pi_{\theta'}(\cdot\mid s)\right)
=\frac12d^\top Fd+o(\|d\|^2),\qquad\theta'=\theta+d

$F$ は報酬を含まない。MLPの記事と同じく、出力のFisherとヤコビアンに分かれる。

Z(s):=\frac{\partial\hat f(s;\theta)}{\partial\theta}\in\mathbb R^{m\times P},\qquad
F=\mathbb E_{s\sim d^\theta}\left[Z(s)^\top F_z\!\left(\hat f(s;\theta)\right)Z(s)\right]

6. TRPO:KL制約付きの二次モデル

第5節で、真の改善量と一次まで一致するサロゲート $L_\theta$ と、方策の変化を測る $F$ が揃った。$L_\theta$ が一次までしか一致しないので、どこまで進んでよいかはまだ決まっていない。

6.1 単調改善の下界

サロゲートと真の改善量の差は、KLの最大値で抑えられる(Schulman et al., 2015)。

\bar A:=\max_{s,a}\left|A^\theta(s,a)\right|,\qquad
D_{\max}(\theta,\theta'):=\max_sD_{\mathrm{KL}}\!\left(\pi_\theta(\cdot\mid s)\middle\|\pi_{\theta'}(\cdot\mid s)\right)
J(\theta')\ge M_\theta(\theta'):=J(\theta)+L_\theta(\theta')-\frac{4\gamma\bar A}{(1-\gamma)^2}D_{\max}(\theta,\theta')

$M_\theta$ を最大化する反復は、$J$ を減らさない。

M_\theta(\theta)=J(\theta),\qquad
\theta_{k+1}=\arg\max_{\theta'}M_{\theta_k}(\theta')
\quad\Longrightarrow\quad
J(\theta_{k+1})\ge M_{\theta_k}(\theta_{k+1})\ge M_{\theta_k}(\theta_k)=J(\theta_k)

$\gamma\to1$ で係数 $4\gamma\bar A/(1-\gamma)^2$ が発散し、ステップが小さくなりすぎる。TRPOはペナルティを平均KLの制約に置き換える。ここから先は下界を動機にした実用上の問題で、単調改善の保証は引き継がない。

\max_{\theta'}L_\theta(\theta')\quad\text{s.t.}\quad\bar D(\theta,\theta')\le\delta

6.2 制約付き二次近似の解

目的を一次、制約を二次で近似する。

\max_d\ g^\top d\quad\text{s.t.}\quad\frac12d^\top Fd\le\delta,\qquad F\succ0,\ g\ne0
\begin{aligned}
g-\nu Fd&=0,\qquad\nu>0,\\
d&=\nu^{-1}F^{-1}g,\\
\frac12d^\top Fd=\frac{g^\top F^{-1}g}{2\nu^2}=\delta
\quad&\Longrightarrow\quad\nu^{-1}=\sqrt{\frac{2\delta}{g^\top F^{-1}g}}
\end{aligned}
\Delta\theta_{\mathrm{TRPO}}=\sqrt{\frac{2\delta}{g^\top F^{-1}g}}\,F^{-1}g

方向は自然勾配と同じで、ステップ長をKLの半径 $\delta$ で決める。ニュートン法の $H_L$ と違い、$F$ は常に半正定値で、5.2節のような符号の反転が起きない。

g^\top\Delta\theta_{\mathrm{TRPO}}=\sqrt{2\delta\,g^\top F^{-1}g}>0

6.3 行列を作らない求解

MLPの記事のHessian-free法と同じく、$F$ を作らずに $Fx=g$ を共役勾配法で解く。標本の状態 $s_1,\ldots,s_n$ で $F$ を近似した行列を $\bar F$ とする。

\bar F:=\frac1n\sum_{i=1}^nF_{s_i},\qquad
\bar Fv=\frac1n\sum_{i=1}^nZ(s_i)^\top F_z\left(Z(s_i)v\right)

7. PPO:クリップによる更新の抑制

TRPOは更新のたびに $\bar F^{-1}g$ を解く。PPOは、方策を近くに保つという動機をTRPOと共有し、それを制約ではなく目的関数の形で追う。$F$ は使わず、同じ標本を一次法で繰り返し使う(Schulman et al., 2017)。

7.1 確率比とKLペナルティ

$\pi_{\theta_{\mathrm{old}}}$ で集めた標本 $(s_t,a_t)$ の平均を $\hat{\mathbb E}_t$ と書く。$\hat A_t$ は $A^{\theta_{\mathrm{old}}}(s_t,a_t)$ の推定値で、価値関数のTD誤差から作る(Schulman et al., 2016)。

\rho_t(\theta):=\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_{\mathrm{old}}}(a_t\mid s_t)},\qquad
L^{\mathrm{CPI}}(\theta):=\hat{\mathbb E}_t\left[\rho_t(\theta)\hat A_t\right]

KLをペナルティとして目的に加える形もある。

L^{\mathrm{KLPEN}}(\theta):=\hat{\mathbb E}_t\left[\rho_t(\theta)\hat A_t-\beta\,D_{\mathrm{KL}}\!\left(\pi_{\theta_{\mathrm{old}}}(\cdot\mid s_t)\middle\|\pi_\theta(\cdot\mid s_t)\right)\right]

サロゲートの項は一次で、KLの項は二次で近似する。目的側の二次項 $H_L$ は捨てる。この二次モデルの最大点は、$\eta=1/\beta$、$\lambda=0$ の自然勾配である。

\hat g:=\nabla_\theta L^{\mathrm{CPI}}(\theta_{\mathrm{old}}),\qquad
L^{\mathrm{KLPEN}}(\theta_{\mathrm{old}}+d)\approx\underbrace{\hat g^\top d}_{\text{サロゲートの一次}}-\underbrace{\frac\beta2d^\top\bar Fd}_{\text{KLの二次}}
\quad\Longrightarrow\quad d^*=\frac1\beta\bar F^{-1}\hat g

PPOはこの $d^*$ を解かず、$L^{\mathrm{KLPEN}}$ を一次法で最大化する。論文の実験では、次のクリップ目的のほうが良い結果を出した。

7.2 クリップ目的

L^{\mathrm{CLIP}}(\theta):=\hat{\mathbb E}_t\left[\min\left(\rho_t(\theta)\hat A_t,\ \operatorname{clip}\left(\rho_t(\theta),1-\varepsilon,1+\varepsilon\right)\hat A_t\right)\right],\qquad\varepsilon>0

標本1個の項は、$A$ の符号で片側だけを切る形になる。

\ell^{\mathrm{CLIP}}(\rho,A):=\min\left(\rho A,\ \operatorname{clip}(\rho,1-\varepsilon,1+\varepsilon)A\right)=
\begin{cases}
A\min(\rho,1+\varepsilon)&A\ge0,\\
A\max(\rho,1-\varepsilon)&A<0
\end{cases}
\nabla_\theta L^{\mathrm{CLIP}}=\hat{\mathbb E}_t\left[m_t\,\hat A_t\,\rho_t\,\nabla_\theta\log\pi_\theta(a_t\mid s_t)\right],\qquad
m_t:=\mathbf1\left[\hat A_t>0,\ \rho_t<1+\varepsilon\right]+\mathbf1\left[\hat A_t<0,\ \rho_t>1-\varepsilon\right]

目的を改善する向きに比が $[1-\varepsilon,1+\varepsilon]$ を出た標本だけ、勾配が0になる。悪化する向きに出た標本の勾配は残る。

\ell^{\mathrm{CLIP}}(\rho,A)\le\rho A\quad\Longrightarrow\quad L^{\mathrm{CLIP}}(\theta)\le L^{\mathrm{CPI}}(\theta)
\theta=\theta_{\mathrm{old}}\quad\Longrightarrow\quad\rho_t=1,\quad m_t=1\ \ (\hat A_t\ne0),\quad
\nabla_\theta L^{\mathrm{CLIP}}=\nabla_\theta L^{\mathrm{CPI}}=\hat g

各反復の最初の勾配は方策勾配そのものである。$m_t=0$ の標本は自分の比を押し出さなくなるが、$\theta$ は全標本で共有されるので、他の標本の勾配で $\rho_t$ は区間の外へ動きうる。クリップは比を区間に制約せず、KLの信頼領域も課さない(Wang et al., 2020; Engstrom et al., 2020)。

7.3 クリップ区間とKL楕円

クリップ区間の幅 $\varepsilon$ とKLの半径 $\delta$ の対応を、線形化して調べる。対象は、全標本の比が区間内にあるという条件を満たす $d$ の集合である。PPOの反復はこの集合を制約として扱わないので、以下は $\varepsilon$ の大きさの目安を与える。

$\theta=\theta_{\mathrm{old}}+d$ で比を一次近似する。

\psi_t:=\nabla_\theta\log\pi_{\theta_{\mathrm{old}}}(a_t\mid s_t),\qquad
\rho_t(\theta_{\mathrm{old}}+d)=1+\psi_t^\top d+O(\|d\|^2)

KLは比の二乗平均で近似できる。$a\sim\pi_{\theta_{\mathrm{old}}}$ なので $\mathbb E[\rho-1]=0$ が厳密に成り立つ。

\begin{aligned}
D_{\mathrm{KL}}\!\left(\pi_{\theta_{\mathrm{old}}}(\cdot\mid s)\middle\|\pi_\theta(\cdot\mid s)\right)
&=\mathbb E_{a\sim\pi_{\theta_{\mathrm{old}}}}\left[-\log\rho\right],\\
-\log\rho&=-(\rho-1)+\frac12(\rho-1)^2+O(|\rho-1|^3),\\
D_{\mathrm{KL}}&=\frac12\mathbb E_{a\sim\pi_{\theta_{\mathrm{old}}}}\left[(\rho-1)^2\right]+O\!\left(\mathbb E|\rho-1|^3\right)
\end{aligned}

$a_t$ は $\pi_{\theta_{\mathrm{old}}}$ からの標本なので、$\hat F$ は6.3節の $\bar F$ のモンテカルロ推定になる。

\hat F:=\hat{\mathbb E}_t\left[\psi_t\psi_t^\top\right],\qquad
\mathbb E_{a_t\sim\pi_{\theta_{\mathrm{old}}}(\cdot\mid s_t)}\left[\hat F\right]=\hat{\mathbb E}_t\left[F_{s_t}\right]=\bar F

線形化した比が全標本で区間内にあれば、標本Fisherによる二次形式は $\varepsilon^2/2$ 以下になる。

|\psi_t^\top d|\le\varepsilon\ \ (\forall t)
\quad\Longrightarrow\quad
\frac12d^\top\hat Fd=\frac12\hat{\mathbb E}_t\left[(\psi_t^\top d)^2\right]\le\frac{\varepsilon^2}2
\underbrace{\left\{d:\ |\psi_t^\top d|\le\varepsilon\ (\forall t)\right\}}_{\text{線形化した比が全標本で区間内}}
\subseteq
\underbrace{\left\{d:\ \tfrac12d^\top\hat Fd\le\tfrac{\varepsilon^2}2\right\}}_{\text{標本FisherによるKL楕円}}

TRPOは楕円を制約として $g^\top d$ を最大化する。PPOは左の集合を制約にせず、改善側に区間を出た標本の勾配を止めるだけである。両者が共有するのは、方策を近くに保つという動機である。

7.4 更新式の比較

ニュートン法、自然勾配、TRPOの三つは、二次モデルの最大点として導いた。TRPOの方向は $\lambda_F=0$ の自然勾配と厳密に一致し、ステップ長だけが違う。PPOは二次モデルから導いた更新ではなく、別の目的関数を一次法で最大化する。

\begin{aligned}
\Delta\theta_{\mathrm{Newton}}&=-(H_L-\lambda_HI)^{-1}g,&&H_L-\lambda_HI\prec0,\\
\Delta\theta_{\mathrm{Natural}}&=\eta(F+\lambda_FI)^{-1}g,&&\lambda_F>0,\\
\Delta\theta_{\mathrm{TRPO}}&=\sqrt{\frac{2\delta}{g^\top F^{-1}g}}\,F^{-1}g,&&\text{共役勾配法},\\
\Delta\theta_{\mathrm{PPO}}&:\ \max_\theta L^{\mathrm{CLIP}}(\theta)\ \text{を一次法で},&&\text{比の区間で勾配を止める}
\end{aligned}

まとめ

\begin{aligned}
g&=\frac1{1-\gamma}\mathbb E_\theta\left[A^\theta\psi_\theta\right],&
H_L&=\frac1{1-\gamma}\mathbb E_\theta\left[A^\theta\left(\nabla_\theta^2\log\pi_\theta+\psi_\theta\psi_\theta^\top\right)\right],\\
F&=\mathbb E_\theta\left[\psi_\theta\psi_\theta^\top\right]\succeq0,&
\bar D(\theta,\theta+d)&=\frac12d^\top Fd+o(\|d\|^2),\\
\Delta\theta_{\mathrm{TRPO}}&\propto F^{-1}g,&
|\rho_t-1|\le\varepsilon\ (\forall t)\ &\overset{\text{一次}}{\Longrightarrow}\ \frac12d^\top\hat Fd\le\frac{\varepsilon^2}2
\end{aligned}

参考文献

  • Engstrom, L., Ilyas, A., Santurkar, S., Tsipras, D., Janoos, F., Rudolph, L., & Madry, A. (2020). Implementation matters in deep policy gradients: A case study on PPO and TRPO. International Conference on Learning Representations. arXiv:2005.12729
  • Kakade, S. M. (2001). A natural policy gradient. Advances in Neural Information Processing Systems 14.
  • Kakade, S., & Langford, J. (2002). Approximately optimal approximate reinforcement learning. Proceedings of the 19th International Conference on Machine Learning, 267–274.
  • Schulman, J., Levine, S., Abbeel, P., Jordan, M., & Moritz, P. (2015). Trust region policy optimization. Proceedings of the 32nd International Conference on Machine Learning, PMLR 37, 1889–1897. arXiv:1502.05477
  • Schulman, J., Moritz, P., Levine, S., Jordan, M., & Abbeel, P. (2016). High-dimensional continuous control using generalized advantage estimation. International Conference on Learning Representations. arXiv:1506.02438
  • Schulman, J., Wolski, F., Dhariwal, P., Radford, A., & Klimov, O. (2017). Proximal policy optimization algorithms. arXiv:1707.06347
  • Sutton, R. S., McAllester, D., Singh, S., & Mansour, Y. (2000). Policy gradient methods for reinforcement learning with function approximation. Advances in Neural Information Processing Systems 12, 1057–1063.
  • Wang, Y., He, H., & Tan, X. (2020). Truly proximal policy optimization. Proceedings of the 35th Conference on Uncertainty in Artificial Intelligence, PMLR 115, 113–122. arXiv:1903.07940
0
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?