PPO(Proximal Policy Optimization)の理論と実装

Actor-Critic法は方策勾配の分散を低減しましたが、まだ重要な問題が残っています。方策のパラメータを大きく更新すると、方策が急激に変わり、学習が不安定になったり、性能が崩壊したりすることがあるのです。

「方策を更新するとき、一度に大きく変えすぎないように制限できないだろうか?」

この問いに対する実用的な回答がPPO(Proximal Policy Optimization) です。2017年にOpenAIのシュルマンらによって提案されたPPOは、方策の更新幅を制限する「クリッピング」手法により、安定した学習を簡潔に実現しました。PPOはその実装の簡便さと安定性から、現在最も広く使われている深層強化学習アルゴリズムの一つです。ChatGPTのRLHF(人間フィードバックからの強化学習)にもPPOが使われています。

PPOを理解すると、以下のような応用が開けます。

  • RLHF: 大規模言語モデルの人間フィードバックによる微調整(ChatGPT, Claude)
  • ロボット制御: MuJoCoなどのシミュレーション環境でのロボット制御学習
  • ゲームAI: OpenAI FiveやDota 2での大規模ゲームAI
  • 自動運転: 安定した方策改善が必要な制御タスク

本記事の内容

  • 信頼領域法(TRPO)の動機と問題点
  • PPOのクリッピング目的関数の導出
  • PPOの完全なアルゴリズム
  • PyTorchによるCartPole環境での実装
  • ハイパーパラメータの影響

前提知識

この記事を読む前に、以下の記事を読んでおくと理解が深まります。

信頼領域法の動機

方策更新の不安定性

方策勾配法で $\theta \leftarrow \theta + \alpha \nabla_\theta J(\theta)$ と更新するとき、学習率 $\alpha$ が大きすぎると方策が急激に変化します。強化学習では、方策が変わるとデータの分布(軌道の分布)も変わるため、一度の大きな更新が連鎖的に学習を崩壊させることがあります。

この不安定性を直感的に理解するためのアナロジーを紹介します。崖の近くを歩いている場面を想像してください。教師あり学習では、地面は平坦で安定しています。大股で歩いても、つまずいてもすぐに立て直せます。しかし強化学習では、地面自体があなたの歩き方に応じて変形します。大股で踏み出すと、足元の地面が予期せぬ方向に傾き、バランスを崩して崖から落ちてしまう可能性があります。小さな一歩ずつ慎重に進めば、地面の変化に対応しながら安全に歩けます。

教師あり学習では、データの分布は固定されているため、パラメータの大きな更新は「一時的にロスが増えるだけ」で済みます。しかし強化学習では、データが方策に依存するため、パラメータの大きな更新は「データの分布自体が変わる」という影響をもたらし、回復が困難な性能崩壊を引き起こす可能性があります。

具体的には、ある状態 $s$ でたまたま高い報酬が得られた行動 $a$ があった場合、方策勾配は $a$ の確率を上げるように作用します。この更新が大きすぎると、他の状態での行動選択も大きく変わってしまい、全体的な性能が低下することがあります。低下した方策で収集した新しいデータはさらに質が悪く、負のスパイラルに陥ることがあります。

TRPO(Trust Region Policy Optimization)

TRPO(Schulman et al., 2015)は、方策の更新をKLダイバージェンスで制限することで安定化しました。

$$ \max_\theta \, \mathbb{E}\left[\frac{\pi_\theta(a|s)}{\pi_{\theta_{\text{old}}}(a|s)} \hat{A}(s, a)\right] \quad \text{s.t.} \quad \mathbb{E}\left[D_{\text{KL}}(\pi_{\theta_{\text{old}}}(\cdot|s) \| \pi_\theta(\cdot|s))\right] \leq \delta $$

$\pi_{\theta_{\text{old}}}$ は更新前の方策、$\hat{A}$ はアドバンテージ推定値、$\delta$ は信頼領域の大きさです。

TRPOの最適化問題を直感的に理解しましょう。この定式化は「方策を改善したいが、一度に大きく変えてはいけない」という要求を数学的に表現しています。目的関数 $\mathbb{E}[r_t(\theta)\hat{A}_t]$ は「方策をどれだけ改善できるか」を表し、制約 $D_{\text{KL}} \leq \delta$ は「方策の変化幅」を制限します。

KLダイバージェンスは2つの確率分布の「距離」を測る指標で、$\pi_{\theta_{\text{old}}}$ と $\pi_\theta$ がどれだけ異なるかを定量化します。$D_{\text{KL}} = 0$ は2つの方策が同一であることを意味し、$\delta$ が小さいほど保守的な更新になります。

TRPOは理論的に優れていますが、制約付き最適化(KLダイバージェンスの制約)を解くために共役勾配法やフィッシャー情報行列とベクトルの積の計算が必要であり、実装が複雑です。具体的には、ラグランジュ乗数法でKL制約を処理し、2次近似のもとで共役勾配法を用いてパラメータの更新方向を計算し、さらにライン探索でステップサイズを調整するという多段階の手続きが必要です。PPOはこの複雑さを解消する代替手法として提案されました。

PPOの基本アイデアは単純です。「KLダイバージェンスの制約」という間接的な方法で方策の変化を制限する代わりに、目的関数自体にクリッピングを組み込んで直接制限してしまおう、というものです。

PPOのクリッピング目的関数

重要度サンプリング

PPOの目的関数を理解するために、まず重要度比率(importance ratio)を定義します。

$$ \begin{equation} r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{\text{old}}}(a_t|s_t)} \end{equation} $$

$r_t = 1$ は更新前後で方策が同じことを意味し、$r_t > 1$ は行動 $a_t$ の確率が増加、$r_t < 1$ は減少したことを意味します。

たとえば、ある状態で「右に動く」確率が古い方策で $\pi_{\theta_{\text{old}}}(a|s) = 0.3$ であり、新しい方策で $\pi_\theta(a|s) = 0.6$ になったとすると、$r_t = 0.6 / 0.3 = 2.0$ です。この行動の確率が2倍に増えたことを意味します。

重要度比率を使うと、方策勾配の目的関数は以下のように書けます(代理目的関数)。

$$ L^{\text{CPI}}(\theta) = \mathbb{E}_t\left[r_t(\theta) \hat{A}_t\right] $$

なぜ重要度比率が必要なのかを説明します。経験データは古い方策 $\pi_{\theta_{\text{old}}}$ のもとで収集されていますが、更新したい方策は $\pi_\theta$ です。異なる分布から集めたデータを使って別の分布の期待値を計算するテクニックが重要度サンプリング(importance sampling)です。古い方策のもとでの期待値を新しい方策のもとでの期待値に変換するために、確率の比 $r_t$ を掛けるのです。

$$ \mathbb{E}_{\pi_\theta}[\hat{A}_t] = \mathbb{E}_{\pi_{\theta_{\text{old}}}}\left[\frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{\text{old}}}(a_t|s_t)} \hat{A}_t\right] = \mathbb{E}_{\pi_{\theta_{\text{old}}}}\left[r_t(\theta) \hat{A}_t\right] $$

この式は Conservative Policy Iteration(CPI)の目的関数であり、TRPOの制約なしバージョンです。この目的関数を制限なしに最大化すると、$r_t$ が極端な値(非常に大きい、または非常に小さい)になり、方策が大幅に変化してしまいます。たとえば $r_t = 100$(確率が100倍)のような状況は、重要度サンプリングの推定値を不安定にし、勾配の分散を爆発させます。

クリッピング手法

PPOは、重要度比率 $r_t$ を区間 $[1-\epsilon, 1+\epsilon]$ にクリッピングすることで、方策の変化を制限します。

$$ \begin{equation} L^{\text{CLIP}}(\theta) = \mathbb{E}_t\left[\min\left(r_t(\theta) \hat{A}_t, \, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) \hat{A}_t\right)\right] \end{equation} $$

ここで $\epsilon$ はクリッピングパラメータ(通常 $\epsilon = 0.2$)です。

$\min$ の役割を理解する: アドバンテージ $\hat{A}_t$ の符号に応じて、クリッピングの効果は以下のように変わります。

$\hat{A}_t > 0$(良い行動)のとき: 行動の確率を上げたい。$r_t$ が大きくなる(確率が増加する)方向に更新したいが、$r_t > 1+\epsilon$ になるとクリッピングされる。$\min$ により、過度な確率増加が抑制されます。

$\hat{A}_t < 0$(悪い行動)のとき: 行動の確率を下げたい。$r_t$ が小さくなる(確率が減少する)方向に更新したいが、$r_t < 1-\epsilon$ になるとクリッピングされる。$\min$ により、過度な確率減少が抑制されます。

いずれの場合も、方策の変化が $\epsilon$ の範囲内に収まるよう制限されるのです。

具体的な数値計算例:クリッピングの効果

$\epsilon = 0.2$ として、具体的な数値でクリッピングの動作を確認しましょう。

ケース1: $\hat{A}_t = +2$(良い行動)、$r_t = 1.5$ の場合

  • クリッピングなし: $r_t \hat{A}_t = 1.5 \times 2 = 3.0$
  • クリッピング後: $\text{clip}(1.5, 0.8, 1.2) \times 2 = 1.2 \times 2 = 2.4$
  • $\min(3.0, 2.4) = 2.4$ → クリッピングが適用される

$r_t = 1.5$ は方策が大きく変化している(確率が1.5倍に増加)ため、クリッピングによって勾配が抑制されます。

ケース2: $\hat{A}_t = +2$(良い行動)、$r_t = 1.1$ の場合

  • クリッピングなし: $1.1 \times 2 = 2.2$
  • クリッピング後: $\text{clip}(1.1, 0.8, 1.2) \times 2 = 1.1 \times 2 = 2.2$
  • $\min(2.2, 2.2) = 2.2$ → クリッピングは適用されない

$r_t = 1.1$ は方策の変化が小さいため、通常の勾配がそのまま使われます。

ケース3: $\hat{A}_t = -1$(悪い行動)、$r_t = 0.5$ の場合

  • クリッピングなし: $0.5 \times (-1) = -0.5$
  • クリッピング後: $\text{clip}(0.5, 0.8, 1.2) \times (-1) = 0.8 \times (-1) = -0.8$
  • $\min(-0.5, -0.8) = -0.8$ → クリッピングが適用される

$r_t = 0.5$ は確率が半分に下がっていることを意味しますが、クリッピングにより $r_t = 0.8$ に制限されます。これにより、悪い行動の確率を急激に下げすぎることが防止されます。

これらの例からわかるように、PPOのクリッピングは「方策の改善を促すが、一歩の大きさを制限する」という効果を持ちます。TRPOのKL制約とは異なり、目的関数レベルで直接制限を行うため、制約付き最適化の複雑さが不要になります。

GAE(Generalized Advantage Estimation)

PPOの性能を支えるもう一つの重要な要素がGAE(Generalized Advantage Estimation)です。アドバンテージ関数 $A(s, a) = Q(s, a) – V(s)$ の推定には、バイアスと分散のトレードオフがあります。

1ステップのTD誤差を $\delta_t = r_t + \gamma V(s_{t+1}) – V(s_t)$ とすると、GAEは

$$ \hat{A}_t^{\text{GAE}} = \sum_{l=0}^{\infty} (\gamma \lambda)^l \delta_{t+l} $$

で定義されます。パラメータ $\lambda \in [0, 1]$ がバイアスと分散のバランスを制御します。

  • $\lambda = 0$: $\hat{A}_t = \delta_t$(1ステップTD誤差)。バイアスが大きいが分散が小さい
  • $\lambda = 1$: $\hat{A}_t = \sum_{l=0}^{\infty} \gamma^l \delta_{t+l}$(モンテカルロ推定に近い)。バイアスが小さいが分散が大きい

実用上は $\lambda = 0.95$ が広く使われています。この値は、数ステップ先までの情報を取り入れつつ、遠い未来の不確実な推定の影響を抑えるバランスが良いためです。

GAEを展開して計算過程を確認しましょう。$\lambda = 0.95$、$\gamma = 0.99$ として、3ステップ分の計算を示します。

$$ \hat{A}_t = \delta_t + (\gamma\lambda)\delta_{t+1} + (\gamma\lambda)^2\delta_{t+2} + \cdots $$

$$ = \delta_t + 0.99 \times 0.95 \times \delta_{t+1} + (0.99 \times 0.95)^2 \times \delta_{t+2} + \cdots $$

$$ = \delta_t + 0.9405 \, \delta_{t+1} + 0.8845 \, \delta_{t+2} + \cdots $$

各ステップのTD誤差に対する重み $(\gamma\lambda)^l$ は指数関数的に減衰するため、直近のステップの情報が最も重視されます。

PPOの完全な損失関数

PPOの完全な損失関数は、クリッピング項、価値関数の損失、エントロピーボーナスの3つで構成されます。

$$ \begin{equation} \mathcal{L}(\theta) = -L^{\text{CLIP}}(\theta) + c_1 \mathcal{L}^{\text{VF}}(\theta) – c_2 \, H(\pi_\theta) \end{equation} $$

$c_1$ は価値損失の係数(通常0.5)、$c_2$ はエントロピー係数(通常0.01)です。

各項の役割を説明します。

第1項: クリッピング損失 $-L^{\text{CLIP}}$: 方策の改善を行う主要な損失です。符号が負なのは、損失関数を最小化するフレームワーク(PyTorch等)で使うためです。$L^{\text{CLIP}}$ を最大化することは $-L^{\text{CLIP}}$ を最小化することと同じです。

第2項: 価値関数損失 $\mathcal{L}^{\text{VF}}$: Criticネットワークの出力 $V_\theta(s)$ が実際の収益に近づくように学習させます。具体的には $\mathcal{L}^{\text{VF}} = (V_\theta(s_t) – R_t)^2$ です。正確な価値推定はGAEの品質に直結するため、方策の学習にも間接的に影響します。

第3項: エントロピーボーナス $H(\pi_\theta)$: 方策のエントロピー $H = -\sum_a \pi(a|s)\log\pi(a|s)$ を高く保つことで、行動の多様性を維持し、早期の局所最適への収束を防ぎます。たとえば、ある状態で左右の行動確率が $(0.5, 0.5)$ ならエントロピーは最大($\log 2$)ですが、$(0.99, 0.01)$ ならエントロピーはほぼゼロです。エントロピーボーナスがないと、方策が早い段階で特定の行動に偏り、より良い行動を発見できなくなることがあります。

これら3つの損失のバランスは $c_1$ と $c_2$ で調整します。通常の設定($c_1 = 0.5$, $c_2 = 0.01$)では、方策改善が主で、価値学習がそれを補助し、エントロピーが穏やかな探索を促すという構造になっています。

それでは、この理論をPyTorchで実装してみましょう。

PyTorchによるPPOの実装

import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.distributions import Categorical
import numpy as np
import matplotlib.pyplot as plt
import gymnasium as gym

torch.manual_seed(42)

# --- PPO Actor-Critic ネットワーク ---
class PPOActorCritic(nn.Module):
    def __init__(self, state_dim, action_dim, hidden_dim=64):
        super().__init__()
        self.actor = nn.Sequential(
            nn.Linear(state_dim, hidden_dim),
            nn.Tanh(),
            nn.Linear(hidden_dim, hidden_dim),
            nn.Tanh(),
            nn.Linear(hidden_dim, action_dim),
        )
        self.critic = nn.Sequential(
            nn.Linear(state_dim, hidden_dim),
            nn.Tanh(),
            nn.Linear(hidden_dim, hidden_dim),
            nn.Tanh(),
            nn.Linear(hidden_dim, 1),
        )

    def get_action(self, state):
        state = torch.FloatTensor(state).unsqueeze(0)
        logits = self.actor(state)
        dist = Categorical(logits=logits)
        action = dist.sample()
        return action.item(), dist.log_prob(action).item()

    def evaluate(self, states, actions):
        logits = self.actor(states)
        dist = Categorical(logits=logits)
        log_probs = dist.log_prob(actions)
        entropy = dist.entropy()
        values = self.critic(states).squeeze(-1)
        return log_probs, entropy, values

# --- GAEの計算 ---
def compute_gae(rewards, values, dones, next_value, gamma=0.99, lam=0.95):
    advantages = []
    gae = 0
    values_list = list(values) + [next_value]

    for t in reversed(range(len(rewards))):
        mask = 1.0 - dones[t]
        delta = rewards[t] + gamma * values_list[t+1] * mask - values_list[t]
        gae = delta + gamma * lam * mask * gae
        advantages.insert(0, gae)

    advantages = torch.tensor(advantages, dtype=torch.float32)
    returns = advantages + torch.tensor(values, dtype=torch.float32)
    return advantages, returns

# --- PPOの学習 ---
env = gym.make('CartPole-v1')
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.n

model = PPOActorCritic(state_dim, action_dim)
optimizer = optim.Adam(model.parameters(), lr=3e-4)

# ハイパーパラメータ
clip_epsilon = 0.2
gamma = 0.99
lam = 0.95
entropy_coeff = 0.01
value_coeff = 0.5
n_steps = 2048
n_epochs = 10  # ミニバッチ更新の繰り返し回数
mini_batch_size = 64
n_updates = 200

all_rewards = []
current_reward = 0
state, _ = env.reset(seed=42)

for update in range(n_updates):
    # --- 経験の収集 ---
    states, actions, rewards, log_probs_old, values, dones = (
        [], [], [], [], [], [])

    for step in range(n_steps):
        action, log_prob = model.get_action(state)
        with torch.no_grad():
            value = model.critic(
                torch.FloatTensor(state).unsqueeze(0)).item()

        next_state, reward, terminated, truncated, _ = env.step(action)
        done = terminated or truncated

        states.append(state)
        actions.append(action)
        rewards.append(reward)
        log_probs_old.append(log_prob)
        values.append(value)
        dones.append(float(done))

        current_reward += reward
        state = next_state

        if done:
            all_rewards.append(current_reward)
            current_reward = 0
            state, _ = env.reset()

    # 次状態の価値
    with torch.no_grad():
        next_value = model.critic(
            torch.FloatTensor(state).unsqueeze(0)).item()

    # GAEの計算
    advantages, returns = compute_gae(rewards, values, dones,
                                       next_value, gamma, lam)
    advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)

    # テンソルに変換
    states_t = torch.FloatTensor(np.array(states))
    actions_t = torch.LongTensor(actions)
    log_probs_old_t = torch.FloatTensor(log_probs_old)

    # --- ミニバッチ更新 ---
    n_samples = len(states)
    for epoch in range(n_epochs):
        indices = np.random.permutation(n_samples)

        for start in range(0, n_samples, mini_batch_size):
            end = start + mini_batch_size
            mb_idx = indices[start:end]

            mb_states = states_t[mb_idx]
            mb_actions = actions_t[mb_idx]
            mb_advantages = advantages[mb_idx]
            mb_returns = returns[mb_idx]
            mb_log_probs_old = log_probs_old_t[mb_idx]

            # 現在の方策で評価
            log_probs_new, entropy, values_new = model.evaluate(
                mb_states, mb_actions)

            # 重要度比率
            ratio = torch.exp(log_probs_new - mb_log_probs_old)

            # クリッピング損失
            surr1 = ratio * mb_advantages
            surr2 = (torch.clamp(ratio, 1-clip_epsilon, 1+clip_epsilon)
                     * mb_advantages)
            actor_loss = -torch.min(surr1, surr2).mean()

            # 価値損失
            critic_loss = F.mse_loss(values_new, mb_returns.detach())

            # エントロピーボーナス
            entropy_loss = -entropy.mean()

            # 合計損失
            loss = (actor_loss + value_coeff * critic_loss
                    + entropy_coeff * entropy_loss)

            optimizer.zero_grad()
            loss.backward()
            nn.utils.clip_grad_norm_(model.parameters(), 0.5)
            optimizer.step()

    if (update + 1) % 10 == 0 and len(all_rewards) > 0:
        recent = all_rewards[-20:] if len(all_rewards) >= 20 else all_rewards
        print(f'Update [{update+1}/{n_updates}] '
              f'Avg Reward: {np.mean(recent):.1f}')

env.close()

PPOの実装の重要なポイントを解説します。まず、2048ステップの経験を収集し(rollout phase)、その後10エポック分のミニバッチ更新を行います。同じデータを複数回使い回す「経験再利用」が可能なのは、重要度比率によるoff-policy補正があるからです。

クリッピング損失のコアは torch.min(surr1, surr2) の部分です。surr1 はクリッピングなしの代理損失、surr2 は重要度比率を $[1-\epsilon, 1+\epsilon]$ にクリッピングした代理損失です。両者の min をとることで、方策の変化が過大になることを防ぎます。

# --- 結果の可視化 ---
fig, axes = plt.subplots(1, 3, figsize=(16, 5))

# (a) 学習曲線
ax = axes[0]
if len(all_rewards) > 10:
    window = min(20, len(all_rewards))
    smooth = np.convolve(all_rewards,
                          np.ones(window)/window, mode='valid')
    ax.plot(smooth, 'b-', linewidth=2)
ax.set_xlabel('Episode', fontsize=12)
ax.set_ylabel('Episode Reward', fontsize=12)
ax.set_title('PPO Learning Curve', fontsize=13)
ax.axhline(500, color='red', linestyle='--', alpha=0.5)
ax.grid(True, alpha=0.3)

# (b) クリッピングの可視化
ax = axes[1]
ratio = np.linspace(0, 2, 200)
eps = 0.2
for adv_sign, color, label in [(1, 'blue', 'A > 0 (good action)'),
                                (-1, 'red', 'A < 0 (bad action)')]:
    surr1 = ratio * adv_sign
    surr2 = np.clip(ratio, 1-eps, 1+eps) * adv_sign
    objective = np.minimum(surr1, surr2)
    ax.plot(ratio, objective, color=color, linewidth=2, label=label)

ax.axvline(1.0, color='gray', linestyle='--', alpha=0.5)
ax.axvline(1-eps, color='gray', linestyle=':', alpha=0.5)
ax.axvline(1+eps, color='gray', linestyle=':', alpha=0.5)
ax.set_xlabel('Importance Ratio $r(\\theta)$', fontsize=12)
ax.set_ylabel('Clipped Objective', fontsize=12)
ax.set_title('PPO Clipping ($\\epsilon=0.2$)', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# (c) 手法比較
ax = axes[2]
methods = ['REINFORCE', 'A2C', 'PPO', 'TRPO']
stability = [2, 5, 8, 9]
simplicity = [9, 7, 8, 3]
sample_eff = [2, 4, 7, 6]

x = np.arange(len(methods))
width = 0.25
ax.bar(x - width, stability, width, label='Stability',
       color='steelblue', alpha=0.8)
ax.bar(x, simplicity, width, label='Simplicity',
       color='coral', alpha=0.8)
ax.bar(x + width, sample_eff, width, label='Sample Efficiency',
       color='green', alpha=0.8)
ax.set_xticks(x)
ax.set_xticklabels(methods, fontsize=11)
ax.set_ylabel('Score (1-10)', fontsize=12)
ax.set_title('Algorithm Comparison', fontsize=13)
ax.legend(fontsize=9)
ax.grid(True, alpha=0.3, axis='y')

plt.tight_layout()
plt.savefig('ppo_result.png', dpi=150, bbox_inches='tight')
plt.show()

この可視化から、PPOの特性を確認できます。

  1. 左図(学習曲線): PPOの報酬推移です。REINFORCEやA2Cと比較して、より速く安定に最大報酬に到達することが期待されます。ミニバッチ更新とクリッピングにより、各更新での改善が確実です

  2. 中央図(クリッピングの可視化): PPOのクリッピング目的関数の形状です。アドバンテージが正(良い行動、青線)の場合、重要度比率が $1+\epsilon$ を超えると目的関数がフラットになり、これ以上確率を上げるインセンティブがなくなります。アドバンテージが負(悪い行動、赤線)の場合、重要度比率が $1-\epsilon$ を下回ると目的関数がフラットになり、これ以上確率を下げるインセンティブがなくなります

  3. 右図(手法比較): 4つのアルゴリズムを安定性、実装の簡便さ、サンプル効率の3軸で比較しています。PPOは3つ全てにおいてバランスが良く、特に安定性と簡便さの両立が特徴です

ハイパーパラメータの影響と実践的なチューニング

PPOの性能はハイパーパラメータの設定に敏感です。各パラメータの影響を理解しておくことは、実用的な問題に適用する際に重要です。

クリッピングパラメータ $\epsilon$

$\epsilon$ はPPO最重要のハイパーパラメータです。一般的に $\epsilon = 0.1 \sim 0.3$ の範囲で設定されます。

  • $\epsilon$ が小さい(例: 0.1): 方策の更新幅が狭く、保守的な学習になります。安定性は高いですが、学習速度が遅くなる可能性があります
  • $\epsilon$ が大きい(例: 0.3): 方策の更新幅が広く、積極的な学習になります。学習が速い場合もありますが、不安定になるリスクが増します
  • $\epsilon = 0.2$: 多くの問題で良好に動作するデフォルト値として広く使われています

ステップ数とエポック数の関係

PPOでは、$n_{\text{steps}}$ ステップの経験を収集し、$n_{\text{epochs}}$ 回のミニバッチ更新を行います。この2つのパラメータの組み合わせが重要です。

$n_{\text{steps}}$ が少なすぎると、アドバンテージの推定が不正確になります。しかし多すぎると、古い方策で集めたデータと現在の方策の乖離が大きくなり、重要度比率の推定が不正確になります。

$n_{\text{epochs}}$ については、同じデータを何度も使い回すとオーバーフィッティングのリスクがあります。重要度比率のクリッピングがこのリスクを軽減しますが、エポック数が多すぎると方策が $[1-\epsilon, 1+\epsilon]$ の範囲の端に張り付いてしまい、学習が停滞することがあります。実用的には $n_{\text{epochs}} = 3 \sim 10$ が推奨されます。

学習率

PPOではAdam最適化器との組み合わせが標準的であり、学習率は $3 \times 10^{-4}$ が広く使われるデフォルト値です。学習の後半でスケジューリング(線形に減衰させる)を行うと、最終的な性能が向上することが報告されています。

アドバンテージの正規化

実装の中で advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8) という正規化を行っています。これは学習の安定化に大きく寄与します。正規化しないと、アドバンテージのスケールがバッチごとに異なり、勾配の大きさが不安定になります。正規化によって、各バッチのアドバンテージの平均がゼロ、標準偏差が1になるため、勾配のスケールが安定します。

PPOの適用上の注意点

PPOは汎用的なアルゴリズムですが、以下の点に注意が必要です。

連続行動空間: 本記事ではカテゴリカル分布(離散行動空間)を使用しましたが、連続行動空間では正規分布を使います。Actorネットワークの出力を平均 $\mu$ と標準偏差 $\sigma$ とし、$\pi_\theta(a|s) = \mathcal{N}(\mu_\theta(s), \sigma_\theta(s)^2)$ とします。

報酬のスケーリング: 報酬のスケールが大きい環境では、報酬の正規化(running mean/stdによるスケーリング)が学習の安定化に効果的です。

並列環境: PPOは複数の環境を並列に動かして経験を収集することで、学習速度を大幅に向上できます。OpenAI Gym の VectorEnv やSubprocVecEnvを使用すると、CPUコア数に比例した高速化が期待できます。

まとめ

本記事では、PPO(Proximal Policy Optimization)の理論と実装について体系的に解説しました。

  • 信頼領域法の動機: 強化学習ではデータ分布が方策に依存するため、方策の大きな更新が連鎖的な性能崩壊を引き起こす。TRPOはKLダイバージェンス制約でこれを解決したが、実装が複雑
  • PPOのクリッピング: 重要度比率 $r_t$ を $[1-\epsilon, 1+\epsilon]$ にクリッピングし、目的関数レベルで直接方策の変化を制限する。TRPOの制約付き最適化を不要にした
  • 重要度サンプリング: 古い方策で収集したデータを新しい方策の評価に使い回すための技法であり、PPOのサンプル効率の基盤
  • GAE: パラメータ $\lambda$ によりバイアスと分散のトレードオフを制御し、アドバンテージ推定の品質を向上させる
  • 3つの損失の組み合わせ: クリッピング損失(方策改善)、価値関数損失(Critic学習)、エントロピーボーナス(探索促進)を一つの損失関数に統合
  • PPOは安定性実装の簡便さサンプル効率のバランスに優れ、RLHF、ロボット制御、ゲームAIなど幅広い分野で標準的に使われている実用的アルゴリズム

次のステップとして、以下の記事も参考にしてください。