Skip to content
AILinkDeepTech
Go back
Reinforcement Learning Medium

Proximal Policy Optimization (PPO) Implementation in PyTorch

Abstract

A PyTorch implementation of Proximal Policy Optimization (PPO) with a clipped surrogate objective, shared actor-critic network, Gaussian policy, and value/entropy losses.

Proximal Policy Optimization (PPO) Implementation in PyTorch

This implementation builds Proximal Policy Optimization (PPO) from scratch. It includes a shared actor-critic network with a Gaussian policy, clipped surrogate objective for stable policy updates, value loss, and entropy regularization for exploration.

import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from torch.distributions import Normal

class PPONetwork(nn.Module):
    def __init__(self, state_dim, action_dim):
        super(PPONetwork, self).__init__()

        # Shared layers between actor and critic
        self.shared = nn.Sequential(
            nn.Linear(state_dim, 64),
            nn.ReLU(),
            nn.Linear(64, 64),
            nn.ReLU()
        )

        # Actor (policy network)
        self.actor_mean = nn.Linear(64, action_dim)
        self.actor_std = nn.Parameter(torch.zeros(action_dim))

        # Critic (value network)
        self.critic = nn.Linear(64, 1)

    def forward(self, state):
        shared_features = self.shared(state)

        action_mean = self.actor_mean(shared_features)
        action_std = torch.exp(self.actor_std)
        
        action_dist = Normal(action_mean, action_std)

        value = self.critic(shared_features)

        return action_dist, value
    
class PPO:
    def __init__(self, state_dim, action_dim, lr=3e-4, gamma=0.99, epsilon=0.2, 
                 value_coef=0.5, entropy_coef=0.01):
        """   
        Args:
            state_dim (int): Dimension of state space
            action_dim (int): Dimension of action space
            lr (float): Learning rate
            gamma (float): Discount factor
            epsilon (float): PPO clipping parameter
            value_coef (float): Value loss coefficient
            entropy_coef (float): Entropy coefficient for exploration
        """
        self.network = PPONetwork(state_dim, action_dim)
        self.optimizer = optim.Adam(self.network.parameters(), lr=lr)
        
        self.gamma = gamma
        self.epsilon = epsilon
        self.value_coef = value_coef
        self.entropy_coef = entropy_coef

    def get_action(self, state):
        state = torch.FloatTensor(state)
        action_dist, value = self.network(state)

        # Sample action 
        action = action_dist.sample()
        log_prob = action_dist.log_prob(action).sum(dim=-1)

        # Convert to scalar values 
        return (action.detach().numpy(),
                float(log_prob.detach().numpy()),  
                float(value.detach().numpy()))   
    
    def update(self, states, actions, old_log_probs, returns, advantages):
        """ 
        Args:
            states (torch.Tensor): Batch of states
            actions (torch.Tensor): Batch of actions
            old_log_probs (torch.Tensor): Batch of log probabilities from old policy
            returns (torch.Tensor): Batch of returns
            advantages (torch.Tensor): Batch of advantages
        """
        states = torch.FloatTensor(states)
        actions = torch.FloatTensor(actions)
        old_log_probs = torch.FloatTensor(old_log_probs)
        returns = torch.FloatTensor(returns)
        advantages = torch.FloatTensor(advantages)

        action_dist, values = self.network(states)

        curr_log_probs = action_dist.log_prob(actions).sum(dim=-1)
        
        # Calculate probability ratio
        ratio = torch.exp(curr_log_probs - old_log_probs)

        # Calculate surrogate objectives
        surr1 = ratio * advantages
        surr2 = torch.clamp(ratio, 1 - self.epsilon, 1 + self.epsilon) * advantages
        
        policy_loss = -torch.min(surr1, surr2).mean()

        value_loss = 0.5 * ((values - returns) ** 2).mean()

        entropy = action_dist.entropy().mean()

        # total loss
        total_loss = (policy_loss + 
                     self.value_coef * value_loss - 
                     self.entropy_coef * entropy)
        
        self.optimizer.zero_grad()
        total_loss.backward()
        self.optimizer.step()

        return (policy_loss.item(), 
                value_loss.item(), 
                entropy.item())
    
def test_ppo():
    state_dim = 4
    action_dim = 2
    ppo = PPO(state_dim, action_dim)

    # action sampling
    state = np.random.rand(state_dim)
    action, log_prob, value = ppo.get_action(state)
    assert action.shape == (action_dim,)
    assert isinstance(log_prob, float), f"log_prob is {type(log_prob)}, expected float"
    assert isinstance(value, float), f"value is {type(value)}, expected float"
    print("passed: initialization and forward pass")

    # Policy update
    batch_size = 32
    states = np.random.rand(batch_size, state_dim)
    actions = np.random.rand(batch_size, action_dim)
    old_log_probs = np.random.rand(batch_size)
    returns = np.random.rand(batch_size)
    advantages = np.random.rand(batch_size)
    
    policy_loss, value_loss, entropy = ppo.update(
        states, actions, old_log_probs, returns, advantages
    )
    
    assert isinstance(policy_loss, float)
    assert isinstance(value_loss, float)
    assert isinstance(entropy, float)
    print("passed: Policy update")

    # Check if losses are reasonable
    assert not np.isnan(policy_loss)
    assert not np.isnan(value_loss)
    assert entropy > -1e6 and entropy < 1e6
    print("passed: Loss values are reasonable")

    return "All tests passed!"

if __name__ == "__main__":
    test_result = test_ppo()
    print(test_result)


Cite this Explanation

@article{ailinkdeeptech2026ppoalgo,
  title={Proximal Policy Optimization (PPO) Implementation in PyTorch},
  author={AILinkDeepTech},
  journal={AILinkDeepTech Algorithm Explanations},
  year={2026},
  url={https://ailinkdeeptech.com/research/ppo_algo}
}

Related Explanations