Chuyển đến nội dung chính

第 7 課:PPO — 近端策略優化

信任域方法。 PPO-Clip 物鏡。廣義優勢估計(GAE)。全面實施 PPO。比較 PPO、TRPO 和 A3C。

🧠 人工智慧與機器學習 — 第 6 課 第 7 課:PPO — 近端策略優化

強化學習:從基礎到高級

第 2 部分:深度強化學習 — 神經網路與 RL 的結合

亞洲開發網

簡介

**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

特點聚苯醚TRPOA3C
簡單✅ 簡單❌ 複雜✅ 簡單
性能高高中
穩定性✅✅❌
可並行化✅✅✅
預設選擇✅沒有沒有

總結

概念描述
PPO 夾透過剪裁限制政策變更
蓋伊平衡偏差-方差優勢估計
熵紅利鼓勵探索
多個時代多次重複使用收集的數據