Skip to content
AILinkDeepTech
Go back
Reinforcement Learning Medium

Group Relative Policy Optimization (GRPO) Implementation in PyTorch

Abstract

A PyTorch implementation of Group Relative Policy Optimization (GRPO): a Group-relative RL approach that partitions sorted trajectories into groups, weights group-relative advantages, and applies clipped surrogate updates with PPO-style ratios.

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}")


Cite this Explanation

@article{ailinkdeeptech2025grpoalgo,
  title={Group Relative Policy Optimization (GRPO) Implementation in PyTorch},
  author={AILinkDeepTech},
  journal={AILinkDeepTech Algorithm Explanations},
  year={2025},
  url={https://ailinkdeeptech.com/research/grpo_algo}
}

Related Explanations