GRPO and Verifiable Rewards — Experiments on Arithmetic#
Validates the core conclusions of grpo.tex:
Verifiable rewards overturn systematic errors: the SFT teacher “cannot carry”; the base model is confidently wrong on carry prompts (carry accuracy 0); rule-scored RLVR lifts overall accuracy from 0.30 to 1.00 in about 20 iterations;
A rule-checked reward: β=0, no KL anchoring, 600 long iterations — accuracy tops out and stays; KL holds at about 1 nat. Contrast Chapter 12 (learned reward, β=0 collapse) and Chapter 13 (DPO’s loosening soft anchor): when the reward is the true objective, the anchor becomes optional;
The death of cold start: when the teacher is never correct on carries, the policy’s sampling success rate ≈ 0 — GRPO’s groups are all-same-score, advantages identically zero, the learning signal vanishes entirely and accuracy stays at 0.30.
Task: two-digit addition. The prompt is d d + d d = (zero-padded); complete 3 digits; reward = rule-verified correctness (0/1). The policy is a small Transformer isomorphic to the previous two chapters (~120k parameters); the SFT teacher gives the correct answer with probability 0.3 and with 0.7 the digit-wise no-carry systematic error — on carry prompts the wrong answer is the data mode, confidently wrong under greedy decoding.
GRPO: sample G=8 completions per prompt; the group-standardized advantage \((r_i-\bar{r})/\mathrm{std}(r)\) replaces the critic; PPO-clip updates. On this task a global batch baseline and the group baseline are both kinds of baselines (this chapter does not contrast them); the group baseline’s decisive difference shows in Experiment 3’s zero-signal group fraction.
Output figures:
fig1_rlvr_overturns.pdffig2_unhackable_reward.pdffig3_cold_start.pdf
Estimated runtime: about 10 minutes on GPU; about 40 minutes on CPU (3 seeds).
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import matplotlib
import matplotlib.pyplot as plt
matplotlib.rcParams.update({
'font.family': 'serif',
'font.serif': ['Times New Roman', 'DejaVu Serif'],
'pdf.fonttype': 42,
'ps.fonttype': 42,
'axes.labelsize': 11,
'xtick.labelsize': 9,
'ytick.labelsize': 9,
'legend.fontsize': 9,
})
BLUE = '#2166AC'
RED = '#D6604D'
GRAY = '#808080'
FIGSIZE = (7.2, 4.8)
OUTDIR = '.'
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
DIGITS = 10
PLUS = 10
EQ = 11
BOS = 12
VOCAB = 13
PROMPT_LEN = 6 # d d + d d =
ANS_LEN = 3 # zero-padded sum 000..198
D_MODEL = 64
MAX_LEN = 1 + PROMPT_LEN + ANS_LEN
SEEDS = (42, 43, 44)
def set_seed(seed):
np.random.seed(seed)
torch.manual_seed(seed)
def make_prompts(n):
a = torch.randint(0, 100, (n,), device=DEVICE)
b = torch.randint(0, 100, (n,), device=DEVICE)
p = torch.stack([a // 10, a % 10, torch.full_like(a, PLUS),
b // 10, b % 10, torch.full_like(a, EQ)], dim=1)
return p, a, b
def answer_digits(sums):
return torch.stack([sums // 100, (sums // 10) % 10, sums % 10], dim=1)
def no_carry_sum(a, b):
"""The teacher's systematic error: digit-wise addition without carrying."""
units = (a % 10 + b % 10) % 10
tens = (a // 10 + b // 10) % 10
return tens * 10 + units
def has_carry(a, b):
units_carry = a % 10 + b % 10 >= 10
tens_carry = a // 10 + b // 10 + (a % 10 + b % 10) // 10 >= 10
return units_carry | tens_carry
def verify(completions, a, b):
"""The verifiable reward: exact-match against the true sum (0/1)."""
return (completions == answer_digits(a + b)).all(dim=1).float()
def style_axes(ax):
ax.spines['top'].set_visible(False)
ax.spines['right'].set_visible(False)
ax.grid(True, linestyle='--', linewidth=0.5, alpha=0.7)
print(f'Device: {DEVICE}')
Device: cuda
class TinyLM(nn.Module):
"""Prompt-conditioned causal Transformer (~120k params)."""
def __init__(self):
super().__init__()
self.emb = nn.Embedding(VOCAB + 1, D_MODEL)
self.pos = nn.Embedding(MAX_LEN, D_MODEL)
layer = nn.TransformerEncoderLayer(D_MODEL, 4, 128, batch_first=True,
dropout=0.0, norm_first=True)
self.encoder = nn.TransformerEncoder(layer, 2, enable_nested_tensor=False)
self.head = nn.Linear(D_MODEL, VOCAB)
mask = torch.triu(torch.full((MAX_LEN, MAX_LEN), float('-inf')), 1)
self.register_buffer('mask', mask)
def logits(self, inp):
L = inp.size(1)
h = self.emb(inp) + self.pos(torch.arange(L, device=inp.device))
h = self.encoder(h, mask=self.mask[:L, :L])
return self.head(h)
@torch.no_grad()
def generate(self, prompts, greedy=False):
n = prompts.size(0)
inp = torch.cat([torch.full((n, 1), BOS, dtype=torch.long, device=DEVICE),
prompts], dim=1)
outs = []
for _ in range(ANS_LEN):
logits = self.logits(inp)[:, -1]
if greedy:
nxt = logits.argmax(dim=-1, keepdim=True)
else:
nxt = torch.multinomial(F.softmax(logits, dim=-1), 1)
outs.append(nxt)
inp = torch.cat([inp, nxt], dim=1)
return torch.cat(outs, dim=1)
def completion_log_probs(self, prompts, completions):
n = prompts.size(0)
inp = torch.cat([torch.full((n, 1), BOS, dtype=torch.long, device=DEVICE),
prompts, completions[:, :-1]], dim=1)
logits = self.logits(inp)[:, PROMPT_LEN:]
return torch.log_softmax(logits, dim=-1).gather(
2, completions.unsqueeze(2)).squeeze(2)
def train_sft(q_correct, steps=1200, batch=256, lr=3e-3):
"""SFT on a flawed teacher: correct with prob q, else the no-carry answer."""
lm = TinyLM().to(DEVICE)
opt = torch.optim.Adam(lm.parameters(), lr=lr)
for _ in range(steps):
prompts, a, b = make_prompts(batch)
use_correct = torch.rand(batch, device=DEVICE) < q_correct
sums = torch.where(use_correct, a + b, no_carry_sum(a, b))
loss = -lm.completion_log_probs(prompts, answer_digits(sums)).mean()
opt.zero_grad(); loss.backward(); opt.step()
return lm
@torch.no_grad()
def eval_accuracy(lm, n=2048):
"""Greedy accuracy, overall and on carry prompts."""
prompts, a, b = make_prompts(n)
acc = verify(lm.generate(prompts, greedy=True), a, b)
carry = has_carry(a, b)
return acc.mean().item(), acc[carry].mean().item()
def grpo_finetune(sft, iters=300, n_prompts=32, group=8, k_epochs=2,
mb=64, clip=0.2, lr=1e-4, beta=0.0, eval_every=5):
"""GRPO: group-relative advantages from the verifiable reward, PPO-clip update.
No critic, no reward model. beta=0 by default -- with a verifiable
reward the KL anchor is optional (Figure 2).
"""
policy = TinyLM().to(DEVICE)
policy.load_state_dict(sft.state_dict())
ref = TinyLM().to(DEVICE)
ref.load_state_dict(sft.state_dict())
ref.eval()
opt = torch.optim.Adam(policy.parameters(), lr=lr)
history = {'it': [], 'acc': [], 'acc_carry': [], 'reward': [], 'kl': [],
'zero_groups': []}
total = n_prompts * group
for it in range(1, iters + 1):
with torch.no_grad():
prompts, a, b = make_prompts(n_prompts)
prompts = prompts.repeat_interleave(group, dim=0)
a = a.repeat_interleave(group)
b = b.repeat_interleave(group)
comp = policy.generate(prompts)
r = verify(comp, a, b)
rg = r.view(n_prompts, group)
adv = ((rg - rg.mean(dim=1, keepdim=True)) /
(rg.std(dim=1, keepdim=True) + 1e-4)).view(-1)
zero_frac = (rg.std(dim=1) < 1e-6).float().mean().item()
lp_old = policy.completion_log_probs(prompts, comp)
lp_ref = ref.completion_log_probs(prompts, comp)
kl_seq = (lp_old - lp_ref).sum(dim=1)
if beta > 0:
adv = adv - beta * kl_seq
for _ in range(k_epochs):
perm = torch.randperm(total, device=DEVICE)
for s in range(0, total, mb):
idx = perm[s:s + mb]
lp_new = policy.completion_log_probs(prompts[idx], comp[idx])
ratio = torch.exp((lp_new - lp_old[idx]).sum(dim=1))
s1 = ratio * adv[idx]
s2 = torch.clamp(ratio, 1 - clip, 1 + clip) * adv[idx]
loss = -torch.min(s1, s2).mean()
opt.zero_grad(); loss.backward()
nn.utils.clip_grad_norm_(policy.parameters(), 1.0)
opt.step()
if it % eval_every == 0:
acc, acc_carry = eval_accuracy(policy)
history['it'].append(it)
history['acc'].append(acc)
history['acc_carry'].append(acc_carry)
history['reward'].append(r.mean().item())
history['kl'].append(kl_seq.mean().item())
history['zero_groups'].append(zero_frac)
return history
print('=== Weak SFT (teacher correct with prob 0.3), per seed ===')
runs = {}
for seed in SEEDS:
set_seed(seed)
weak = train_sft(q_correct=0.3)
acc, acc_carry = eval_accuracy(weak)
prompts, a, b = make_prompts(2048)
sampled = verify(weak.generate(prompts), a, b).mean().item()
runs[seed] = {'weak': weak, 'base_acc': acc, 'base_carry': acc_carry}
print(f' seed {seed}: greedy acc {acc:.3f} (carry {acc_carry:.3f}) | '
f'sampled acc {sampled:.3f}')
=== Weak SFT (teacher correct with prob 0.3), per seed ===
seed 42: greedy acc 0.303 (carry 0.000) | sampled acc 0.524
seed 43: greedy acc 0.317 (carry 0.036) | sampled acc 0.488
seed 44: greedy acc 0.326 (carry 0.040) | sampled acc 0.518
Figure 1 — Verifiable rewards overturn a systematic error#
The base model’s error is not noise but algorithmic: the teacher cannot carry; on carry prompts the wrong answer is 70% of the data, and greedy decoding is confidently all wrong (carry accuracy ≈ 0). Learned rewards are helpless here — trained on the same biased data, they would only absorb the bias. The rule scorer is different: correct is correct. GRPO uses the ~30% correct rate surviving in the samples as signal and lifts carry accuracy from 0 to 1.0 in about 20 iterations (3 seeds, shading \(\pm 1\sigma\)).
print('=== Experiment 1: GRPO from the weak SFT ===')
for seed in SEEDS:
set_seed(seed + 2000)
runs[seed]['weak_hist'] = grpo_finetune(runs[seed]['weak'])
h = runs[seed]['weak_hist']
print(f" seed {seed}: acc {runs[seed]['base_acc']:.2f} -> {h['acc'][-1]:.2f} | "
f"carry {runs[seed]['base_carry']:.2f} -> {h['acc_carry'][-1]:.2f}")
=== Experiment 1: GRPO from the weak SFT ===
seed 42: acc 0.30 -> 1.00 | carry 0.00 -> 1.00
seed 43: acc 0.32 -> 1.00 | carry 0.04 -> 1.00
seed 44: acc 0.33 -> 1.00 | carry 0.04 -> 1.00
def stack(key, hist_key='weak_hist'):
return np.array([runs[s][hist_key][key] for s in SEEDS])
fig, ax = plt.subplots(figsize=FIGSIZE)
its = np.array(runs[SEEDS[0]]['weak_hist']['it'])
for key, color, label in [('acc', BLUE, 'overall accuracy'),
('acc_carry', RED, 'carry-prompt accuracy')]:
curve = stack(key)
ax.plot(its, curve.mean(0), color=color, linewidth=1.5, label=label)
ax.fill_between(its, curve.mean(0) - curve.std(0),
curve.mean(0) + curve.std(0), color=color, alpha=0.15)
base_acc = np.mean([runs[s]['base_acc'] for s in SEEDS])
base_carry = np.mean([runs[s]['base_carry'] for s in SEEDS])
ax.axhline(base_acc, color=BLUE, linestyle=':', linewidth=1.0,
label=f'SFT overall ({base_acc:.2f})')
ax.axhline(base_carry, color=RED, linestyle=':', linewidth=1.0,
label=f'SFT carry ({base_carry:.2f})')
style_axes(ax)
ax.set_xlabel('GRPO iteration')
ax.set_ylabel('Greedy accuracy')
ax.set_ylim(-0.05, 1.05)
ax.set_title('Fig 1. Verifiable reward overturns a systematic error\n'
'(3 seeds, mean $\\pm$ std)', pad=8)
ax.legend(loc='lower right')
fig.tight_layout()
fig.savefig(f'{OUTDIR}/fig1_rlvr_overturns.pdf', bbox_inches='tight')
plt.show()
print('Saved fig1_rlvr_overturns.pdf')
Saved fig1_rlvr_overturns.pdf
Figure 2 — Rule-checked reward: β=0 does not collapse over long training#
In Chapter 12, β=0 PPO pierces the learned reward within 60 iterations; in Chapter 13, DPO’s soft anchor loosens with training. Here likewise β=0 for 600 iterations — nothing happens: accuracy tops out and stays put, KL steady at about 1 nat. The reason is simple: the reward and the true objective are the same function; with no gap between proxy and truth, harder optimization only locks correct answers in more firmly. “The reward gets hacked” is impossible by definition under RLVR — the remaining risk lies only in the verifier itself (answer-format bypasses, incomplete tests), an engineering problem, no longer a learning problem (3 seeds, shading \(\pm 1\sigma\)).
print('=== Experiment 2: beta = 0, 600 iterations, nothing collapses ===')
for seed in SEEDS:
set_seed(seed + 3000)
runs[seed]['long_hist'] = grpo_finetune(runs[seed]['weak'], iters=600)
h = runs[seed]['long_hist']
print(f" seed {seed}: final acc {h['acc'][-1]:.3f} | "
f"final KL {h['kl'][-1]:.2f} | max KL {max(h['kl']):.2f}")
=== Experiment 2: beta = 0, 600 iterations, nothing collapses ===
seed 42: final acc 1.000 | final KL 0.95 | max KL 1.32
seed 43: final acc 1.000 | final KL 0.86 | max KL 1.39
seed 44: final acc 0.999 | final KL 0.84 | max KL 1.15
fig, ax = plt.subplots(figsize=FIGSIZE)
its = np.array(runs[SEEDS[0]]['long_hist']['it'])
acc = stack('acc', 'long_hist')
ax.plot(its, acc.mean(0), color=BLUE, linewidth=1.5, label='accuracy (= true objective)')
ax.fill_between(its, acc.mean(0) - acc.std(0), acc.mean(0) + acc.std(0),
color=BLUE, alpha=0.15)
ax.set_xlabel('GRPO iteration')
ax.set_ylabel('Greedy accuracy', color=BLUE)
ax.set_ylim(0.25, 1.05)
ax2 = ax.twinx()
kl = stack('kl', 'long_hist')
ax2.plot(its, kl.mean(0), color=GRAY, linewidth=1.5, label='KL to SFT reference')
ax2.fill_between(its, kl.mean(0) - kl.std(0), kl.mean(0) + kl.std(0),
color=GRAY, alpha=0.15)
ax2.set_ylabel('KL (nats)', color=GRAY)
style_axes(ax)
ax2.spines['top'].set_visible(False)
ax.set_title('Fig 2. $\\beta=0$, 600 iterations: nothing collapses\n'
'(3 seeds, mean $\\pm$ std)', pad=8)
h1, l1 = ax.get_legend_handles_labels()
h2, l2 = ax2.get_legend_handles_labels()
ax.legend(h1 + h2, l1 + l2, loc='center right')
fig.tight_layout()
fig.savefig(f'{OUTDIR}/fig2_unhackable_reward.pdf', bbox_inches='tight')
plt.show()
print('Saved fig2_unhackable_reward.pdf')
Saved fig2_unhackable_reward.pdf
Figure 3 — The death of cold start: all-same groups carry no gradient#
GRPO’s advantage comes from within-group comparison: \((r_i-\bar{r})/\mathrm{std}(r)\). Comparison requires within-group differences. Switch the teacher to never correct on carries (q=0): the cold-start model’s sampling success rate on carry prompts ≈ 0, so every group is either all-correct (easy) or all-wrong (hard) — zero within-group standard deviation, advantages and gradients identically zero. Right panel, the zero-signal group fraction: cold start constant at 100%, learning never begins; the weak-teacher start climbs from about three in ten to near 100% — a different meaning: everything mastered, “graduated”. One statistic, two destinies. This is why reasoning-model training needs cold-start SFT data (3 seeds, shading \(\pm 1\sigma\)).
print('=== Experiment 3: cold start (teacher never correct on carries) ===')
for seed in SEEDS:
set_seed(seed + 4000)
cold = train_sft(q_correct=0.0)
acc, acc_carry = eval_accuracy(cold)
set_seed(seed + 5000)
runs[seed]['cold_hist'] = grpo_finetune(cold)
h = runs[seed]['cold_hist']
print(f' seed {seed}: cold SFT acc {acc:.3f} (carry {acc_carry:.3f}) -> '
f"GRPO final {h['acc'][-1]:.3f} | zero-signal groups "
f"{h['zero_groups'][-1]:.2f}")
=== Experiment 3: cold start (teacher never correct on carries) ===
seed 42: cold SFT acc 0.294 (carry 0.000) -> GRPO final 0.320 | zero-signal groups 1.00
seed 43: cold SFT acc 0.311 (carry 0.000) -> GRPO final 0.292 | zero-signal groups 1.00
seed 44: cold SFT acc 0.305 (carry 0.000) -> GRPO final 0.296 | zero-signal groups 1.00
fig, axes = plt.subplots(1, 2, figsize=(10, 4.2))
its = np.array(runs[SEEDS[0]]['weak_hist']['it'])
ax = axes[0]
for hist_key, color, label in [('weak_hist', BLUE, 'weak teacher (q=0.3)'),
('cold_hist', RED, 'cold start (q=0)')]:
acc = stack('acc', hist_key)
ax.plot(its, acc.mean(0), color=color, linewidth=1.5, label=label)
ax.fill_between(its, acc.mean(0) - acc.std(0), acc.mean(0) + acc.std(0),
color=color, alpha=0.15)
style_axes(ax)
ax.set_xlabel('GRPO iteration')
ax.set_ylabel('Greedy accuracy')
ax.set_ylim(-0.05, 1.05)
ax.set_title('Accuracy', pad=8)
ax.legend(loc='center right')
ax = axes[1]
for hist_key, color, label in [('weak_hist', BLUE, 'weak teacher (q=0.3)'),
('cold_hist', RED, 'cold start (q=0)')]:
z = stack('zero_groups', hist_key)
ax.plot(its, z.mean(0), color=color, linewidth=1.5, label=label)
ax.fill_between(its, z.mean(0) - z.std(0), z.mean(0) + z.std(0),
color=color, alpha=0.15)
style_axes(ax)
ax.set_xlabel('GRPO iteration')
ax.set_ylabel('Zero-signal group fraction')
ax.set_ylim(-0.05, 1.05)
ax.set_title('Groups with zero advantage', pad=8)
ax.legend(loc='center right')
fig.suptitle('Fig 3. Cold start: all-same groups carry no gradient '
'(3 seeds, mean $\\pm$ std)', y=1.02)
fig.tight_layout()
fig.savefig(f'{OUTDIR}/fig3_cold_start.pdf', bbox_inches='tight')
plt.show()
print('Saved fig3_cold_start.pdf')
Saved fig3_cold_start.pdf
Summary#
Verifiable rewards remove the root disease: all the trouble of Chapters 12 and 13 came from “the reward is a proxy learned from finitely covered data”; when task correctness is rule-decidable, proxy and truth merge on this task, so the proxy cannot be hacked — β=0 does not collapse over long training, the anchor becomes optional;
GRPO replaces the critic with within-group comparison: one prompt, one group, standardized advantages — normalized automatically by prompt difficulty, no value network; the group mechanism’s distinctive behavior shows in the zero-signal groups (this chapter does not contrast a global batch baseline);
The new bottleneck is signal, not reward: verifiable rewards are usually sparse (right/wrong); when a group is all-same-score the advantage is identically zero — at cold start the learning signal vanishes entirely. Fixes come from starting points and data: cold-start SFT, curricula, mixed difficulty to preserve within-group variance;
The trilogy closes: RLHF taught us “rewards can be learned but will be hacked”; DPO taught us “skipping the component does not skip the coverage”; RLVR showed that a rule-checked reward cannot be hacked on this toy task. That is a teaching contrast, not a ranking of production post-training methods.