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)