# How to Implement Proximal Policy Optimization (PPO) for LLM Fine-Tuning

> Learn how to implement Proximal Policy Optimization PPO for LLM fine-tuning. Stabilize RLHF training with trust regions and advantage estimation.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: tutorial
- Published: 2026-06-11

---

**Proximal Policy Optimization (PPO) aligns language models by clipping policy updates to a trust region, using Generalized Advantage Estimation and a value head to stabilize RLHF training after supervised fine-tuning.**

The `train-llm-from-scratch` repository by FareedKhan-dev implements a complete PPO pipeline for large language model (LLM) alignment. This implementation treats token generation as a sequential decision-making process, where the LLM acts as the policy and a dedicated value head estimates state values for advantage computation.

## Core PPO Components in [`src/post_training/ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/ppo.py)

The PPO implementation is deliberately modularized into discrete, testable functions. Each component handles a specific mathematical operation required by the clipped surrogate objective.

### Generalized Advantage Estimation with `compute_gae`

The `compute_gae` function (lines 24-57 in [`src/post_training/ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/ppo.py)) calculates advantages using GAE with parameters `gamma` (discount factor) and `lam` (smoothing parameter). For language modeling, `gamma` is typically set to `1.0` because the final reward depends on the complete sequence, while `lam=0.95` smooths the advantage estimates across response tokens.

This function consumes reward tensors, current value estimates, and next-step values, returning both `advantages` and `returns` tensors masked to response positions only.

### Advantage Normalization with `whiten`

Before computing the policy loss, the advantage vector undergoes whitening via the `whiten` function (lines 60-66). This normalizes advantages to zero mean and unit variance exclusively on response positions (using `resp_mask`), which stabilizes the policy gradient step and mitigates variance explosion across different batch compositions.

### Clipped Surrogate Loss with `ppo_policy_loss`

The `ppo_policy_loss` function (lines 68-82) implements the core PPO objective: `min(r·A, clip(r)·A)`, where `r` is the probability ratio between new and old policies. The function accepts `clip=0.2` as the default epsilon value, returning both the scalar loss and the fraction of clipped ratios (`clip_frac`) for diagnostic monitoring.

### Value Function Clipping with `ppo_value_loss`

To prevent the value network from over-optimizing during updates, `ppo_value_loss` (lines 84-95) applies a clipped value loss. It computes the squared error between predicted values and GAE returns, but clips the value update using `vf_clip=0.2` to keep the value function from moving too far from its initialization during each epoch.

### KL Divergence Monitoring with `approx_kl`

The `approx_kl` function (lines 98-100) provides a lightweight KL-divergence estimator used strictly for logging. This metric ensures PPO updates remain within a predefined KL budget, acting as a health check that prevents the policy from diverging too rapidly from the reference model.

## Data Flow: From Rollout to Parameter Update

The PPO training loop follows a strict six-stage pipeline orchestrated by [`scripts/train_ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_ppo.py):

1. **Rollout Generation** – The current policy generates responses, storing `logp` (old log-probabilities), `value` (state-value estimates), and `reward` (per-token KL penalty plus terminal task reward) for each token position.

2. **GAE Computation** – The `compute_gae` function processes these tensors alongside a `resp_mask` that flags assistant-generated tokens, producing advantage and return estimates.

3. **Advantage Whitening** – The `whiten` function normalizes advantages using only the masked response positions, ensuring stable gradient magnitudes.

4. **Loss Calculation** – `ppo_policy_loss` computes the clipped surrogate objective, while `ppo_value_loss` handles the value function update. Both respect the response mask to exclude prompt tokens from the loss.

5. **KL Monitoring** – `approx_kl` estimates the divergence between old and new policies, providing a metric for trust region validation without contributing to gradients.

6. **Optimization Step** – The combined loss (policy + 0.5 × value + optional entropy) backpropagates through the policy and value heads, followed by a gradient step and refresh of the old policy copy.

## End-to-End Implementation Example

Below is a minimal training loop demonstrating how the PPO utilities integrate with a transformer model. This assumes you have already generated `old_logp`, `old_value`, and `rewards` from a rollout phase, and that your model includes a value head for state-value prediction.

```python
import torch
from src.post_training.ppo import (
    compute_gae,
    whiten,
    ppo_policy_loss,
    ppo_value_loss,
    approx_kl,
)
from src.post_training.utils import masked_mean

# Rollout tensors: (B, L) batch size x sequence length

old_logp = torch.randn(4, 128)      # Old policy log-probabilities

old_value = torch.randn(4, 128)     # Old value predictions

rewards = torch.randn(4, 128)       # Per-token rewards (KL penalty + final)

resp_mask = torch.randint(0, 2, (4, 128))  # 1 = assistant token, 0 = prompt

# 1. Compute GAE advantages and returns

advantages, returns = compute_gae(
    rewards=rewards,
    values=old_value,
    values_next=old_value,  # Bootstrap from same values for simplicity

    resp_mask=resp_mask,
    gamma=1.0,              # No discount for sequence-level rewards

    lam=0.95,               # GAE smoothing parameter

)

# 2. Whiten advantages (zero mean, unit variance on response tokens only)

advantages = whiten(advantages, resp_mask)

# 3. Forward pass with new policy (hypothetical function)

new_logp, new_value = policy_forward(tokens)  # Your model forward here

# 4. Compute PPO losses

policy_loss, clip_frac = ppo_policy_loss(
    new_logp=new_logp,
    old_logp=old_logp,
    advantages=advantages,
    mask=resp_mask,
    clip=0.2,
)

value_loss = ppo_value_loss(
    new_values=new_value,
    old_values=old_value,
    returns=returns,
    mask=resp_mask,
    vf_clip=0.2,
)

# 5. KL divergence for monitoring (no gradients)

kl_div = approx_kl(new_logp, old_logp, resp_mask)

# 6. Optimization step

total_loss = policy_loss + 0.5 * value_loss
optimizer.zero_grad()
total_loss.backward()
optimizer.step()

print(f"Policy Loss: {policy_loss.item():.4f}")
print(f"Value Loss: {value_loss.item():.4f}")
print(f"Clip Fraction: {clip_frac.item():.4f}")
print(f"Approx KL: {kl_div.item():.4f}")

```

## Training Orchestration and Monitoring

The [`scripts/train_ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_ppo.py) script automates the entire PPO workflow, handling distributed training, checkpointing via `utils.save_stage_ckpt`, and logging. It coordinates with [`src/post_training/rollout.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/rollout.py) for response generation and [`src/post_training/reward_model.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/reward_model.py) for computing the reward signal that drives the PPO objective.

For visualization, [`ui/pages/6_PPO.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/ui/pages/6_PPO.py) provides a Streamlit interface that renders real-time plots of KL divergence, policy loss curves, and reward trajectories, enabling manual inspection of training stability.

## Summary

- **PPO for LLMs** uses a clipped surrogate objective (`ppo_policy_loss`) to limit policy updates, preventing catastrophic forgetting during RLHF.
- **GAE and whitening** (`compute_gae` and `whiten` in [`src/post_training/ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/ppo.py)) stabilize advantage estimation by normalizing only across response tokens.
- **Response masking** is essential—all PPO functions accept a `resp_mask` to ensure gradients flow only through assistant-generated tokens, not user prompts.
- **Value clipping** (`ppo_value_loss`) constrains the value head updates to maintain training stability.
- **The repository** provides end-to-end orchestration through [`scripts/train_ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_ppo.py) and monitoring via [`ui/pages/6_PPO.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/ui/pages/6_PPO.py).

## Frequently Asked Questions

### What distinguishes PPO from standard supervised fine-tuning for LLMs?

Supervised fine-tuning (SFT) optimizes the log-likelihood of reference responses, while PPO treats text generation as a reinforcement learning problem. PPO uses a reward model to score outputs and optimizes a clipped surrogate objective that maximizes expected reward while penalizing large deviations from the reference policy. This allows the model to explore better responses beyond the training distribution while maintaining coherence.

### Why does the implementation require a response mask in all PPO functions?

The response mask (`resp_mask`) ensures that loss calculations and advantage estimates apply only to tokens generated by the assistant, excluding user prompts and padding. In [`src/post_training/ppo.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/post_training/ppo.py), functions like `whiten` and `ppo_policy_loss` use this mask to compute statistics and losses exclusively over actual response positions, preventing the policy from optimizing over irrelevant or user-controlled tokens.

### What are the recommended hyperparameters for PPO when fine-tuning LLMs?

Based on the `train-llm-from-scratch` implementation, use `gamma=1.0` for the discount factor (since language rewards are typically sequence-level), `lam=0.95` for GAE smoothing, and `clip=0.2` for both the policy ratio (`ppo_policy_loss`) and value function (`ppo_value_loss`). These values balance exploration with stability, keeping the KL divergence between updated and reference policies within a reasonable trust region.

### How does the value head architecture differ from the language modeling head?

The value head is a separate linear layer appended to the base transformer ([`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py)) that predicts scalar state-values for each token position. While the language modeling head outputs vocabulary logits for next-token prediction, the value head outputs a single value estimate used in GAE computation. During PPO training, both heads receive gradients, but only the value head is subject to the clipped value loss constraint.