RL Agents Part 3: GRPO, MCTS, and Search-Based Reasoning

← Back to Home

Now that we understand RLVR and reward signal density, we reach the next layer: how do we actually optimize agents to maximize those rewards? This post covers two breakthrough approaches: GRPO (used by DeepSeek-R1) and MCTS (used by AlphaProof and competitive reasoning systems).

Code repo: RL_Agents

Notebook: post04_grpo_mcts.ipynb

1. The RL Optimization Problem

Once you have reward signals (from RLVR), the question becomes: how do you adjust the model's weights to produce higher-reward trajectories?

The classical approach is PPO (Proximal Policy Optimization), which works like this:

  1. Sample trajectories from your current model
  2. Compute advantages (how much better is this trajectory than average?)
  3. Use gradient updates to increase the probability of high-advantage trajectories
  4. Regularize: use KL divergence to keep the updated model close to the original (don't drift too far)

PPO works well, but it has a significant drawback: computing advantages requires a separate critic network - another neural network that predicts expected returns. For large models, this means duplicating the model size, doubling computation and memory.

2. GRPO: Group Relative Policy Optimization

GRPO (Group Relative Policy Optimization) is a clever alternative that sidesteps the critic network entirely. The core insight: instead of comparing each trajectory to a learned baseline, compare outputs within a group.

How GRPO Works:
  1. Sample a batch of outputs: For a single prompt, generate multiple responses (say, 8 candidate solutions).
  2. Score them all: Use your reward model (or verifier) to score each response.
  3. Compute relative advantages: Instead of comparing to an expected value, compute advantages relative to the group mean. High scorer gets positive advantage. Low scorer gets negative advantage.
  4. Update the policy: Increase probability of high-advantage responses, decrease low-advantage responses.

Why is this efficient?

Simple Example:

Suppose you're training a math-solving agent. For the prompt "What is 5 + 3?", you generate 4 candidate answers:

Group mean reward = (1 + 0 + 1 + 0) / 4 = 0.5

Advantages: A gets +0.5, B gets -0.5, C gets +0.5, D gets -0.5

During the gradient update, you increase the probability of generating responses like A and C, and decrease the probability of responses like B and D. This is much cheaper than training a critic network.

3. Monte Carlo Tree Search + LLMs

An orthogonal approach: instead of pure RL fine-tuning, use MCTS (Monte Carlo Tree Search) at inference time to explore the space of reasoning paths.

Classic MCTS in Games:

MCTS is the algorithm that powers AlphaGo. It works by:

  1. Selection: Starting from the root (current state), pick the most promising path using an exploration strategy (UCB1 formula).
  2. Expansion: When you reach a leaf node, add new children (new possible moves).
  3. Simulation: Run a fast simulation from the expanded node to estimate its value.
  4. Backpropagation: Update all nodes along the path with the simulation result.
Applying MCTS to Reasoning:

For agent tasks, the analogy is:

The agent generates candidate next steps (using the language model), MCTS explores them using the value function, and the path with the highest estimated value is executed.

Example: Solving a Math Problem with MCTS

Prompt: "Solve for x: 3x + 5 = 20"

  1. Root: Initial state is the problem statement.
  2. MCTS samples steps:
    • Path A: "Subtract 5 from both sides" → "3x = 15"
    • Path B: "Divide both sides by 3" → "x + 5/3 = 20/3" (wrong approach)
    • Path C: "Subtract 5 from both sides" → "3x = 15" (same as A)
  3. Value function evaluates each: Paths A and C look promising (score 0.9). Path B looks like a dead end (score 0.3).
  4. MCTS explores further: Expands path A/C: "Divide both sides by 3" → "x = 5"
  5. Final step: Verify "x = 5" against the original equation. It works! Done.
Why MCTS for Reasoning?

4. GRPO vs MCTS: When to Use Which?

Aspect GRPO (RL Fine-tuning) MCTS (Inference Search)
When applied During training At inference (test) time
Cost Amortized during training Per-query (can be slow)
How it helps Model learns better policies Explores best solutions at test time
Requires Verifiable reward signal Verifiable reward + value function
Best for Scalable improvement, batch processing Complex problems needing explicit search

The future approach: Use both. Fine-tune with GRPO to improve the base model. Then add MCTS at inference time for extra-hard problems where search helps.

5. The Branching Factor Problem in Language

MCTS is incredibly powerful, but it faces a challenge in language: the branching factor is huge. In Go, each position has ~300 legal moves. In language generation, each token could be any of ~100k vocabulary items, leading to combinatorial explosion.

Solutions:

This is why AlphaProof worked: it combined a strong language model (proposing reasonable steps) with a value function (filtering implausible paths) and beam search (exploring top candidates). Not full MCTS, but the principle is the same.

6. Real-World Impact

The models leading the leaderboards now use some blend of:

DeepSeek-R1's success came largely from this exact stack: GRPO training + test-time inference search. The cost tradeoff is real (inference is slower), but the capability jump is enormous.

7. Complete Code: GRPO and MCTS Demonstrations

The code demonstrates the credit-assignment difference between ORM and PRM on a minimal chain task.

from __future__ import annotations
import numpy as np
import pandas as pd


class ChainMDP:
    """
    5-step chain MDP.
    Action: 0 or 1. Correct = 1 at every step.
    ORM: +1 at end only if all 5 steps correct.
    PRM: +1/5 per correct step (dense).
    """
    N_STEPS     = 5
    CORRECT_ACT = 1

    def run_episode(self, policy: "TabularPolicy") -> tuple[list[int], float, float]:
        actions     = [policy.sample(s) for s in range(self.N_STEPS)]
        n_correct   = sum(a == self.CORRECT_ACT for a in actions)
        orm_reward  = 1.0 if n_correct == self.N_STEPS else 0.0
        prm_reward  = n_correct / self.N_STEPS
        return actions, orm_reward, prm_reward


class TabularPolicy:
    """
    Independent softmax per state.
    Updated via REINFORCE: logit[s][a] += lr * (R - b) * (1[a] - pi(a|s))
    """
    def __init__(self, n_states: int, n_actions: int = 2,
                 lr: float = 0.1, seed: int = 0) -> None:
        self.logits = np.zeros((n_states, n_actions))
        self.lr     = lr
        self.rng    = np.random.default_rng(seed)

    def probs(self, state: int) -> np.ndarray:
        e = np.exp(self.logits[state] - self.logits[state].max())
        return e / e.sum()

    def sample(self, state: int) -> int:
        return int(self.rng.choice(self.logits.shape[1], p=self.probs(state)))

    def update(self, state: int, action: int,
               reward: float, baseline: float) -> None:
        advantage    = reward - baseline
        grad         = -self.probs(state)
        grad[action] += 1.0
        self.logits[state] += self.lr * advantage * grad


def _run_training(reward_mode: str, n_episodes: int = 400, seed: int = 0) -> pd.DataFrame:
    env      = ChainMDP()
    policy   = TabularPolicy(n_states=env.N_STEPS, lr=0.1, seed=seed)
    baseline = 0.5
    alpha_b  = 0.05
    history  = []
    window   = []

    for ep in range(n_episodes):
        actions, orm_r, prm_r = env.run_episode(policy)
        if reward_mode == "orm":
            for s, a in enumerate(actions):
                policy.update(s, a, orm_r, baseline)
            episode_reward = orm_r
        else:
            for s, a in enumerate(actions):
                step_r = 1.0 if a == ChainMDP.CORRECT_ACT else 0.0
                policy.update(s, a, step_r, baseline)
            episode_reward = prm_r

        baseline += alpha_b * (episode_reward - baseline)
        window.append(orm_r)
        if len(window) > 30:
            window.pop(0)

        history.append({
            "episode":          ep + 1,
            "reward_mode":      reward_mode,
            "success":          int(orm_r == 1.0),
            "success_rate_30":  round(float(np.mean(window)), 3),
            "p_correct_step0":  round(float(policy.probs(0)[ChainMDP.CORRECT_ACT]), 3),
        })

    return pd.DataFrame(history)


def train_orm(n_episodes: int = 400, seed: int = 0) -> pd.DataFrame:
    """Train with Outcome Reward (sparse: reward at episode end only)."""
    return _run_training("orm", n_episodes, seed)


def train_prm(n_episodes: int = 400, seed: int = 0) -> pd.DataFrame:
    """Train with Process Reward (dense: reward at every step)."""
    return _run_training("prm", n_episodes, seed)


# Compare ORM vs PRM
orm_df = train_orm()
prm_df = train_prm()
print(f"ORM final success rate (last 30 ep): {orm_df['success_rate_30'].iloc[-1]:.1%}")
print(f"PRM final success rate (last 30 ep): {prm_df['success_rate_30'].iloc[-1]:.1%}")
orm_ep50 = orm_df.loc[orm_df["success_rate_30"] >= 0.50, "episode"].min()
prm_ep50 = prm_df.loc[prm_df["success_rate_30"] >= 0.50, "episode"].min()
print(f"ORM episodes to 50% success: {orm_ep50}")
print(f"PRM episodes to 50% success: {prm_ep50}")
     

When you run this code, you'll see that PRM reaches 50% success in roughly half the episodes that ORM needs. The dense per-step signal gives the policy a clearer gradient at every step rather than waiting for a single terminal reward.

Complete Code Reference

The code above is from the RL_Agents repository:

To run locally:

git clone https://github.com/Pulkit12dhingra/RL_Agents
   cd RL_Agents
   uv sync
   jupyter notebook notebooks/post04_grpo_mcts.ipynb
     

8. Summary

GRPO: Efficient policy optimization that avoids critic networks by using group-relative advantages.

MCTS: Classical search algorithm adapted to language, exploring multiple reasoning paths at inference time.

Combined: A trained model that explores well at inference, achieving state-of-the-art results on complex reasoning.

In the next post, we'll explore online learning: how deployed agents continuously improve by learning from their own trajectories.

Next in the series: Part 4 explores online learning from deployed agent trajectories - STaR and self-improvement loops.
← Part 2 Part 4 →