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

Lesson 7: PPO — Proximal Policy Optimization

Trust region methods. PPO-Clip objective. Generalized Advantage Estimation (GAE). Full PPO implementation. Compare PPO vs TRPO vs A3C.

🧠 AI & ML — Lesson 6 Lesson 7: PPO — Proximal Policy Optimization

Reinforcement Learning: From Basics to Advanced

Part 2: Deep Reinforcement Learning — Neural Networks meet RL

xdev.asia

Introduction

PPO (Proximal Policy Optimization) is the most used algorithm in modern RL — from game AI (OpenAI Five, Dota 2) to RLHF (ChatGPT). PPO is simple, stable, and effective.


1. Motivation: Why do we need PPO?

The original gradient policy has problems:

  • Too small step: Slow
  • Too large step: Policy collapse

PPO limits policy updates to "trust regions" — not changing too much at a time.


2. PPO-Clip Objective

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

In which:

  • $r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{old}}(a_t|s_t)}$ — probability ratio
  • $\hat{A}_t$ — estimated advantage
  • $\epsilon$ — clipping range (usually 0.2)

3. GAE — Generalized Advantage Estimation

$$\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. Full PPO Implementation

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

FeaturesPPOTRPOA3C
Simplicity✅ Simple❌ Complex✅ Simple
PerformanceHighHighMedium
Stability✅✅❌
Parallelizable✅✅✅
Default choice✅NoNo

Summary

ConceptsDescription
PPO-ClipLimit policy changes by clipping
GAEBalanced bias-variance advantage estimation
Entropy bonusEncourage exploration
Multiple epochsReuse collected data many times