Group Relative Policy Optimization (GRPO) Implementation in PyTorch
This implementation builds Group Relative Policy Optimization (GRPO) from scratch. It includes a categorical policy network, a Trajectory container with discount-return computation, a GRPO agent that groups trajectories by relative return and applies PPO-style clipped surrogate updates weighted by group rank, plus train/evaluate helpers and CartPole/Acrobot test cases.
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.distributions import Categorical
import gym
torch.autograd.set_detect_anomaly(True)
class PolicyNetwork(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super(PolicyNetwork, self).__init__()
self.network = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, output_dim)
)
def forward(self, x):
return F.softmax(self.network(x), dim=1)
def get_action(self, state):
state_tensor = torch.FloatTensor(state).unsqueeze(0)
with torch.no_grad():
probs = self.forward(state_tensor)
# Sample action
dist = Categorical(probs)
action = dist.sample()
return action.item(), probs[0, action]
class Trajectory:
def __init__(self):
self.states = []
self.actions = []
self.rewards = []
self.probs = []
self.dones = []
def add(self, state, action, reward, prob, done):
self.states.append(state)
self.actions.append(action)
self.rewards.append(reward)
self.probs.append(prob)
self.dones.append(done)
def get_return(self, gamma):
G = 0
returns = []
for r, d in zip(reversed(self.rewards), reversed(self.dones)):
G = r + gamma * G * (1 - int(d))
returns.insert(0, G)
return np.mean(returns)
def to_tensor(self, gamma=0.99):
states = torch.FloatTensor(np.array(self.states))
actions = torch.LongTensor(np.array(self.actions))
probs = torch.stack([torch.tensor(p) for p in self.probs])
returns = self._compute_returns(gamma)
return states, actions, probs, returns
def _compute_returns(self, gamma=0.99):
returns = []
G = 0
for r, d in zip(reversed(self.rewards), reversed(self.dones)):
G = r + gamma * G * (1 - int(d))
returns.insert(0, G)
return torch.FloatTensor(returns)
class GRPO:
def __init__(self, env, hidden_dim=64, lr_policy=0.001, gamma=0.99, n_groups=3, clip_param=0.2):
self.env = env
self.state_dim = env.observation_space.shape[0]
self.action_dim = env.action_space.n
self.gamma = gamma
self.n_groups = n_groups
self.clip_param = clip_param
self.policy_net = PolicyNetwork(self.state_dim, hidden_dim, self.action_dim)
self.policy_optimizer = optim.Adam(self.policy_net.parameters(), lr=lr_policy)
def collect_trajectories(self, n_trajectories):
trajectories = []
for _ in range(n_trajectories):
traj = Trajectory()
state, _ = self.env.reset()
done = False
while not done:
action, prob = self.policy_net.get_action(state)
next_state, reward, terminated, truncated, _ = self.env.step(action)
done = terminated or truncated
traj.add(state, action, reward, prob, done)
state = next_state
trajectories.append(traj)
return trajectories
def group_trajectories(self, trajectories):
# Sort trajectories by return
traj_with_returns = [(traj, traj.get_return(self.gamma)) for traj in trajectories]
sorted_trajectories = [t for t, _ in sorted(traj_with_returns, key=lambda x: x[1])]
grouped_trajectories = []
group_size = max(1, len(sorted_trajectories) // self.n_groups)
for i in range(0, len(sorted_trajectories), group_size):
group = sorted_trajectories[i:i + group_size]
if len(group) > 0:
grouped_trajectories.append(group)
# Ensure we don't have more than n_groups
while len(grouped_trajectories) > self.n_groups:
if len(grouped_trajectories) >= 2:
grouped_trajectories[-2].extend(grouped_trajectories[-1])
grouped_trajectories.pop()
return grouped_trajectories
def update_policy(self, grouped_trajectories):
"""Update policy using group relative approach"""
for group_idx, group in enumerate(grouped_trajectories):
# Group weight scales with group index
group_weight = (group_idx + 1) / len(grouped_trajectories)
for trajectory in group:
states, actions, old_probs, returns = trajectory.to_tensor(self.gamma)
if len(states) == 0:
continue
current_probs = self.policy_net(states)
dist = Categorical(current_probs)
# Get log probabilities for the actions taken
log_probs = dist.log_prob(actions)
old_log_probs = torch.log(old_probs + 1e-10) # Add small epsilon to avoid log(0)
# Calculate ratios and surrogate losses
ratios = torch.exp(log_probs - old_log_probs)
surr1 = ratios * returns * group_weight
surr2 = torch.clamp(ratios, 1.0 - self.clip_param, 1.0 + self.clip_param) * returns * group_weight
policy_loss = -torch.min(surr1, surr2).mean()
self.policy_optimizer.zero_grad()
policy_loss.backward()
self.policy_optimizer.step()
def train(self, n_episodes, n_trajectories_per_update=10):
"""Train the agent for n_episodes"""
rewards_history = []
for episode in range(n_episodes):
trajectories = self.collect_trajectories(n_trajectories_per_update)
avg_reward = np.mean([sum(traj.rewards) for traj in trajectories])
rewards_history.append(avg_reward)
grouped_trajectories = self.group_trajectories(trajectories)
self.update_policy(grouped_trajectories)
if (episode + 1) % 10 == 0:
print(f"Episode {episode+1}, Average Reward: {avg_reward:.2f}")
return rewards_history
def evaluate(self, n_episodes=10, render=False):
"""Evaluate the agent for n_episodes"""
rewards = []
for _ in range(n_episodes):
state, _ = self.env.reset()
done = False
total_reward = 0
while not done:
if render:
self.env.render()
action, _ = self.policy_net.get_action(state)
state, reward, terminated, truncated, _ = self.env.step(action)
done = terminated or truncated
total_reward += reward
rewards.append(total_reward)
avg_reward = np.mean(rewards)
print(f"Evaluation: Average Reward over {n_episodes} episodes: {avg_reward:.2f}")
return avg_reward
# Test case for CartPole environment
def test_cartpole():
print("Testing GRPO on CartPole-v1...")
env = gym.make('CartPole-v1')
grpo = GRPO(
env=env,
hidden_dim=64,
lr_policy=0.001,
gamma=0.99,
n_groups=3,
clip_param=0.2
)
# Train the agent
rewards = grpo.train(n_episodes=50, n_trajectories_per_update=5)
# Evaluate the agent
avg_reward = grpo.evaluate(n_episodes=10)
return rewards, avg_reward
# Test case for Acrobot environment
def test_acrobot():
print("Testing GRPO on Acrobot-v1...")
env = gym.make('Acrobot-v1')
grpo = GRPO(
env=env,
hidden_dim=128,
lr_policy=0.001,
gamma=0.99,
n_groups=4,
clip_param=0.1
)
# Train the agent
rewards = grpo.train(n_episodes=30, n_trajectories_per_update=8)
# Evaluate the agent
avg_reward = grpo.evaluate(n_episodes=10)
return rewards, avg_reward
if __name__ == "__main__":
cartpole_rewards, cartpole_avg_reward = test_cartpole()
acrobot_rewards, acrobot_avg_reward = test_acrobot()
print("\nResults Summary:")
print(f"CartPole-v1 Average Evaluation Reward: {cartpole_avg_reward:.2f}")
print(f"Acrobot-v1 Average Evaluation Reward: {acrobot_avg_reward:.2f}")