簡介
**PPO(近端策略優化)**是現代強化學習中最常用的演算法-從遊戲 AI(OpenAI Five、Dota 2)到 RLHF (ChatGPT)。 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 vs TRPO vs A3C
| 特點 | 聚苯醚 | TRPO | A3C |
|---|---|---|---|
| 簡單 | ✅ 簡單 | ❌ 複雜 | ✅ 簡單 |
| 性能 | 高 | 高 | 中 |
| 穩定性 | ✅ | ✅ | ❌ |
| 可並行化 | ✅ | ✅ | ✅ |
| 預設選擇 | ✅ | 沒有 | 沒有 |
總結
| 概念 | 描述 |
|---|---|
| PPO 夾 | 透過剪裁限制政策變更 |
| 蓋伊 | 平衡偏差-方差優勢估計 |
| 熵紅利 | 鼓勵探索 |
| 多個時代 | 多次重複使用收集的數據 |