Most explanations of PPO for language models stop at the objective: a clipped ratio, an advantage, a KL penalty. That is enough to understand the idea and nowhere near enough to make it work. Real runs fail on details the equations do not show: a log-prob shifted by one position, padding tokens counted in an average, a reward placed on the wrong token, or a rollout engine whose numbers disagree slightly with the trainer's.
This article is the implementer's view. It assumes you know what PPO is trying to do (if not, the derivation is in the PPO math for LLMs) and walks through one training iteration as code on GPUs: the tensors, the masks, the reward, the update schedule and the numbers to watch. It ends with a run that looks successful and is not, and how the dashboard gives it away.
One iteration, in plain terms
Each iteration takes a batch of prompts and does four things. It generates responses with the current policy. It scores each finished response once with a reward model. It computes, for every generated token, how far the policy has drifted from a frozen reference model, and turns that into a small per-token penalty. Then it runs a few passes of gradient descent that push up the probability of tokens that led to better-than-expected outcomes, while the clip stops any single update from moving the policy too far.
On GPUs this means three or four networks live in memory (policy, value head or critic, reference, reward model), and the work alternates between generation, which is memory-bandwidth bound and dominated by the KV cache, and training, which is compute bound. How to lay those out across devices is covered in the RLHF pipeline on GPUs. Here we focus on what happens to the tensors.
Tensors and masks: where most bugs live
After generation, each example is one row of input_ids with shape [B, T]: prompt tokens, then response tokens, then padding. With left-padded prompts (common for batched generation) every response starts at the same column, the padded prompt width; with right-padded or packed prompts the start differs per row and you must record it per row.
A causal language model's logits at position t predict the token at t+1. So the log-prob of the token actually generated at position t+1 comes from logits at t. Every per-token tensor in PPO (log-probs, KL, rewards, values, advantages) should be aligned to the same shifted grid of length T-1, and the response mask must be built on that grid too. Get this wrong by one and the run still trains, slowly and badly, which is the worst kind of bug.
import torch, torch.nn.functional as F
def token_logprobs(model, input_ids, attn_mask):
"""Log-prob of each actual next token. Output shape [B, T-1], aligned to input_ids[:, 1:]."""
logits = model(input_ids=input_ids, attention_mask=attn_mask).logits[:, :-1, :]
logits = logits.float() # log_softmax in fp32; bf16 here makes ratios noisy
return torch.gather(F.log_softmax(logits, -1), 2, input_ids[:, 1:].unsqueeze(-1)).squeeze(-1)
def response_mask(attn_mask, resp_start):
"""1 for generated, non-pad tokens, aligned like token_logprobs (shifted by one).
resp_start: [B] column index where generation began (the padded prompt width if left-padded)."""
B, T = attn_mask.shape
pos = torch.arange(1, T, device=attn_mask.device).expand(B, -1) # position of the predicted token
return ((pos >= resp_start[:, None]) & attn_mask[:, 1:].bool()).float()Two habits catch most mask bugs. Assert that mask.sum(1) equals the generated length you recorded at rollout time. And compute every mean as (x * mask).sum() / mask.sum(), never x.mean(), because padding otherwise dilutes losses in proportion to how short the responses are.
Per-token rewards: KL everywhere, score at the end
The reward model gives one number per sequence. PPO needs a reward per token. The standard construction gives every response token a reward of -beta * (logp_policy - logp_ref), a per-token estimate of the KL divergence from the reference, and adds the sequence score to the last real response token only. Placing the score on the final token, not spread across all of them, is what lets GAE assign credit backwards through the value function.
Responses that hit the length limit without an end-of-sequence token deserve care. The reward model scores a truncated answer as if it were complete, and the policy can learn that rambling until cut off scores well. A fixed penalty for missing EOS is a common, blunt and effective fix.
def per_token_rewards(score, logp_old, logp_ref, mask, beta=0.05, no_eos=None, eos_penalty=-1.0):
kl = (logp_old - logp_ref) * mask # per-token KL estimate
rewards = -beta * kl
last = mask.sum(1).long() - 1 + mask.argmax(1) # index of last response token
final = score.clone()
if no_eos is not None:
final = torch.where(no_eos, final + eos_penalty, final) # truncated answers pay a penalty
rewards[torch.arange(len(score)), last] += final
return rewards, kl
def gae(rewards, values, mask, gamma=1.0, lam=0.95):
adv = torch.zeros_like(rewards); last = torch.zeros(rewards.size(0), device=rewards.device)
for t in reversed(range(rewards.size(1))):
nxt = values[:, t + 1] if t + 1 < rewards.size(1) else 0.0
nxt = nxt * mask[:, t + 1] if t + 1 < rewards.size(1) else nxt
delta = rewards[:, t] + gamma * nxt - values[:, t]
last = (delta + gamma * lam * last) * mask[:, t]
adv[:, t] = last
returns = adv + values
mean = (adv * mask).sum() / mask.sum()
var = ((adv - mean) ** 2 * mask).sum() / mask.sum()
return (adv - mean) * torch.rsqrt(var + 1e-8) * mask, returns # whiten over real tokens onlyHere gamma=1.0 is the usual choice for LLM PPO because the episode is short and there is no reason to discount the final score. Whitening advantages over real tokens only, across the whole batch, keeps their scale stable from iteration to iteration as rewards drift.
Recompute old log-probs in the trainer
Fast rollouts usually come from an inference engine such as vLLM, which returns log-probs for the tokens it sampled. It is tempting to use those as logp_old. Do not. The engine uses different kernels, may sample from a different precision and batches differently, so its log-probs differ from the trainer's by small amounts. In the ratio exp(logp - logp_old) those small differences look like policy change, inflate the clip fraction and add noise to every update.
Instead, run one no-grad forward pass of the policy in the trainer, with exactly the precision and settings the update will use, and store that as logp_old. The same pass gives you values. A useful check: on the first minibatch of the first epoch, before any optimiser step, the ratio should be 1.0 to within a few thousandths. If it is not, something differs between the two forward passes (dropout left on, a different attention implementation, mismatched padding) and you should fix it before tuning anything.
The update: epochs, minibatches and clipping
The rollout batch is reused for a few epochs, shuffled into minibatches. More reuse squeezes more learning from expensive generations, but each pass moves the policy away from the one that produced the data, and the clip only bounds that drift locally. Two epochs of four minibatches is a reasonable starting point; values above four epochs tend to show rising KL with little gain.
def ppo_update(policy, value_head, opt, batch, epochs=2, minibatches=4, clip=0.2, vclip=0.2, ent_coef=0.0):
N = batch["input_ids"].size(0); stats = []
for epoch in range(epochs):
for idx in torch.randperm(N).chunk(minibatches):
mb = {k: v[idx] for k, v in batch.items()}
m = mb["mask"]; n = m.sum()
out = policy(input_ids=mb["input_ids"], attention_mask=mb["attn"], output_hidden_states=True)
logits = out.logits[:, :-1].float()
logp = torch.gather(F.log_softmax(logits, -1), 2, mb["input_ids"][:, 1:, None]).squeeze(-1)
ratio = torch.exp(logp - mb["logp_old"])
pg = torch.max(-mb["adv"] * ratio, -mb["adv"] * ratio.clamp(1 - clip, 1 + clip))
v = value_head(out.hidden_states[-1][:, :-1]).squeeze(-1)
v_clip = mb["values"] + (v - mb["values"]).clamp(-vclip, vclip)
vf = torch.max((v - mb["returns"]) ** 2, (v_clip - mb["returns"]) ** 2)
ent = -(F.softmax(logits, -1) * F.log_softmax(logits, -1)).sum(-1)
loss = ((pg + 0.5 * vf - ent_coef * ent) * m).sum() / n
opt.zero_grad(); loss.backward()
torch.nn.utils.clip_grad_norm_(policy.parameters(), 1.0); opt.step()
with torch.no_grad():
stats.append(dict(
approx_kl=(((ratio - 1) - torch.log(ratio)) * m).sum().item() / n.item(),
clipfrac=(((ratio - 1).abs() > clip).float() * m).sum().item() / n.item(),
entropy=(ent * m).sum().item() / n.item(), epoch=epoch))
return statsThree details matter. The value loss uses its own clip around the old values, so the critic cannot jump either. The entropy bonus is often zero for LLMs, because a pretrained model starts with plenty of entropy; add it only if entropy collapses. And gradient clipping at a norm of about 1.0 protects against the occasional huge-advantage outlier. With FSDP or ZeRO sharding the same code runs per rank, but statistics must be all-reduced before you log them, or you will chart one rank's view. Sharding itself is covered in FSDP on GPUs.
The dashboard: what healthy looks like
PPO runs fail quietly. The reward line can climb for days while the model gets worse. Log these every iteration, and learn their normal shapes:
| Metric | What it tells you | Rule-of-thumb healthy range |
|---|---|---|
| approx_kl (per update) | How far one round of updates moved the policy | Roughly 0.001 to 0.05; spikes mean the learning rate or epochs are too high |
| clipfrac | Share of tokens where the clip was active | Roughly 0.05 to 0.3; near 0 means tiny steps, above 0.4 means too aggressive |
| first-minibatch ratio | Agreement between rollout and update passes | 1.0 within about 0.005 |
| KL to reference (per sequence) | Total drift from the starting model | Grows slowly; a sudden climb precedes reward hacking |
| value explained variance | Whether the critic predicts returns | Rises towards 0.5 or above; negative means the critic is worse than a constant |
| entropy (per token) | Diversity of the policy's choices | Declines slowly; a collapse means repetitive text |
| response length and EOS rate | Shape of what the model produces | Stable unless you intend it to change |
| reward score, mean and spread | What you are optimising | Rises, then plateaus |
These ranges are working heuristics drawn from common practice, not standards; your model and reward will set their own baseline. Record a baseline for the first few hundred iterations and alert on departures from it.
Worked example: a run that looks great and is not
In this illustrative run, a team fine-tunes a 7B model with PPO against a helpfulness reward model. After 300 iterations the mean score has risen from 0.2 to 1.9. The dashboard tells a different story. Mean response length has gone from 180 to 610 tokens. The EOS rate has fallen from 98% to 71%, so almost a third of answers are being cut off at the length limit. KL to reference climbed slowly until iteration 200, then tripled in 40 iterations. Entropy is down by half.
The diagnosis: the reward model prefers longer answers, a known bias, and it scores truncated answers without noticing they are truncated. The policy found that writing more raises the score. Nothing was broken in the PPO code; the optimiser did exactly what it was asked.
The fixes are layered. Add the missing-EOS penalty from the reward code above. Raise beta from 0.02 to 0.05 to hold the policy closer to the reference. Add length-controlled pairs to the reward model's training data, or normalise the score by a length baseline. Roll back to the iteration-200 checkpoint rather than continuing from 300. And add a gate: every 50 iterations, generate on a fixed held-out prompt set and have a separate judge, not the training reward model, compare against the previous checkpoint. If the training reward rises while the held-out judge does not, stop.
Failure catalogue
- Off-by-one alignment. Log-probs taken from the wrong position. Symptom: learning is slow and noisy for no obvious reason. Fix: one shared alignment helper and a unit test on a toy sequence.
- Padding in averages. Symptom: loss scale depends on batch composition. Fix: masked means everywhere.
- Engine log-probs used as old log-probs. Symptom: clip fraction high from the first minibatch. Fix: recompute in the trainer.
- Critic not learning. Symptom: explained variance near zero or negative. Fix: warm up the value head for a few iterations with the policy frozen, or raise its learning rate.
- Reward hacking. Symptom: score up, held-out quality flat or down, length or format drifting. Fix: higher beta, EOS penalty, a better reward model and a held-out judge.
- Out of memory during generation. The KV cache for long responses at large batch dominates. Fix: cap response length, generate in chunks, or move rollout to a dedicated engine.
- Weights out of sync. The rollout engine keeps the old weights after an update. Symptom: first-minibatch ratio drifts away from 1.0 over time. Fix: verify a checksum after each sync.
If PPO's moving parts are more than your problem needs, offline preference methods remove the rollout loop entirely; the trade-offs are in DPO training on GPUs.
What to do next
- Write one alignment helper for log-probs and masks, with a unit test on a hand-built sequence.
- Replace every
.mean()in your losses with a masked mean. - Recompute old log-probs and values in the trainer, and assert the first-minibatch ratio is 1.0.
- Place the score on the last response token and add a missing-EOS penalty.
- Log approx_kl, clipfrac, KL to reference, explained variance, entropy, length and EOS rate from day one.
- Add a held-out judge evaluation every 50 iterations and stop the run when it disagrees with the training reward.