はじめに
PPO (Proximal Policy Optimization) は、ゲーム AI (OpenAI Five、Dota 2) から RLHF (ChatGPT) に至るまで、最新の RL で最も使用されているアルゴリズムです。 PPO はシンプルで安定しており、効果的です。
1. 動機: なぜ PPO が必要なのでしょうか?
元の勾配ポリシーには次の問題があります。
- ステップが小さすぎる: 遅い
- 大きすぎるステップ: 政策の崩壊
PPO は、ポリシーの更新を「信頼できる領域」に制限し、一度にあまり変更しないようにします。
2. PPO クリップの目的
$$L^{CLIP}(\theta) = \mathbb{E}[\min(r_t(\theta)\hat{A}_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon)\hat{A}_t)]$$
その中で:
- $r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{old}}(a_t|s_t)}$ — 確率比
- $\hat{A}_t$ — 推定アドバンテージ
- $\epsilon$ — クリッピング範囲 (通常は 0.2)
3. GAE — 一般化された利点の推定
$$\hat{A}t^{GAE(\gamma,\lambda)} = \sum{l=0}^{\infty} (\gamma\lambda)^l \delta_{t+l}$$
def compute_gae(rewards, values, dones, gamma=0.99, lam=0.95):
advantages = []
gae = 0
for t in reversed(range(len(rewards))):
if t == len(rewards) - 1:
next_value = 0
else:
next_value = values[t + 1]
delta = rewards[t] + gamma * next_value * (1 - dones[t]) - values[t]
gae = delta + gamma * lam * (1 - dones[t]) * gae
advantages.insert(0, gae)
return torch.FloatTensor(advantages)
4. 完全な PPO の実装
class PPO:
def __init__(self, state_dim, action_dim, lr=3e-4, gamma=0.99,
lam=0.95, clip_eps=0.2, epochs=10, batch_size=64):
self.actor_critic = ActorCritic(state_dim, action_dim)
self.optimizer = torch.optim.Adam(self.actor_critic.parameters(), lr=lr)
self.gamma = gamma
self.lam = lam
self.clip_eps = clip_eps
self.epochs = epochs
self.batch_size = batch_size
def update(self, states, actions, rewards, dones, old_log_probs, values):
advantages = compute_gae(rewards, values, dones, self.gamma, self.lam)
returns = advantages + torch.FloatTensor(values)
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
for _ in range(self.epochs):
probs, new_values = self.actor_critic(states)
dist = Categorical(probs)
new_log_probs = dist.log_prob(actions)
entropy = dist.entropy().mean()
# PPO-Clip
ratio = (new_log_probs - old_log_probs).exp()
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
policy_loss = -torch.min(surr1, surr2).mean()
# Value loss
value_loss = nn.MSELoss()(new_values.squeeze(), returns)
# Total loss
loss = policy_loss + 0.5 * value_loss - 0.01 * entropy
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm_(self.actor_critic.parameters(), 0.5)
self.optimizer.step()
5. PPO 対 TRPO 対 A3C
| 特長 | PPO | TRPO | A3C |
|---|---|---|---|
| シンプルさ | ✅ シンプル | ❌ 複雑 | ✅ シンプル |
| パフォーマンス | 高 | 高 | 中 |
| 安定性 | ✅ | ✅ | ❌ |
| 並列化可能 | ✅ | ✅ | ✅ |
| デフォルトの選択肢 | ✅ | いいえ | いいえ |
概要
| コンセプト | 説明 |
|---|---|
| PPOクリップ | クリッピングによるポリシー変更の制限 |
| ゲイ | バランスの取れたバイアス分散優位性の推定 |
| エントロピーボーナス | 探検を奨励する |
| 複数のエポック | 収集したデータを何度も再利用 |