Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions docs/operators/grpo-loss.md
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,13 @@ loss = masked_mean(policy, completion_mask) + beta * masked_mean(kl, completion_

The Triton op matches the native reference (forward and backward) to `atol=1e-4`.

The per-group variance is computed two-pass, `E[(x - mean)^2]`, in every backend
(native, Triton, distributed). The one-pass form `E[x^2] - E[x]^2` cancels
catastrophically in fp32 when the rewards share a large offset: rewards
`1e4 + [0, 1, 2, 3]` used to yield advantages around `±5e5` instead of `±1.34`.
Advantages are therefore shift-invariant, `adv(r + c) == adv(r)` up to fp32
rounding of `r + c`.

For `DistributedGRPOLossOp`, the reference-equals-policy identity is exact rather
than approximate: with `ref_logits is policy_logits` and `old_logps == logp_policy`
the ratio is `exp(0) = 1` bitwise, so the result is invariant to the clip epsilon,
Expand Down
15 changes: 9 additions & 6 deletions rl_engine/kernels/ops/pytorch/loss/grpo_loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,15 +73,18 @@ def group_advantages(
0, group_id, torch.ones_like(flat_rewards)
)
sums = flat_rewards.new_zeros(num_groups).index_add_(0, group_id, flat_rewards)
sq_sums = flat_rewards.new_zeros(num_groups).index_add_(
0, group_id, flat_rewards * flat_rewards
)

means = sums / counts
variance = (sq_sums / counts) - means * means

# Two-pass variance: centring before squaring keeps the result meaningful
# when the rewards share a large offset, which E[x^2] - E[x]^2 does not.
centered = flat_rewards - means[group_id]
sq_devs = flat_rewards.new_zeros(num_groups).index_add_(
0, group_id, centered * centered
)
variance = sq_devs / counts
stds = variance.clamp_min(eps**2).sqrt()

return (flat_rewards - means[group_id]) / stds[group_id]
return centered / stds[group_id]

@staticmethod
def expand_advantages(
Expand Down
14 changes: 8 additions & 6 deletions rl_engine/kernels/ops/triton/loss/grpo_loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,12 +46,14 @@ def _group_norm_kernel(

count = (end - start).to(tl.float32)
mean = tl.sum(rewards, axis=0) / count
# Population variance (unbiased=False): E[x^2] - E[x]^2. Masked lanes are 0.
sq_mean = tl.sum(rewards * rewards, axis=0) / count
std = tl.sqrt(tl.maximum(sq_mean - mean * mean, 0.0))
std = tl.maximum(std, eps)

adv = (rewards - mean) / std
# Two-pass population variance (unbiased=False): E[(x - mean)^2]. Centring
# before squaring avoids the cancellation E[x^2] - E[x]^2 suffers when the
# rewards share a large offset. Masked lanes must stay 0 after centring.
centered = tl.where(keep, rewards - mean, 0.0)
var = tl.sum(centered * centered, axis=0) / count
std = tl.maximum(tl.sqrt(var), eps)

adv = centered / std
tl.store(adv_ptr + start + offs, adv, mask=keep)


Expand Down
50 changes: 50 additions & 0 deletions tests/test_grpo_loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,27 @@ def _reference_loss(batch, policy_logits, ref_logits, advantages, clip_eps, beta
return policy_loss + beta * kl, policy_loss, kl


def _float64_group_advantages(rewards, group_boundaries, eps=1e-6):
"""Two-pass float64 reference, immune to E[x^2] - E[x]^2 cancellation."""
r = rewards.reshape(-1).double().cpu()
parts = []
for start, end in zip(group_boundaries[:-1], group_boundaries[1:]):
centered = r[start:end] - r[start:end].mean()
std = centered.pow(2).mean().sqrt().clamp_min(eps)
parts.append(centered / std)
return torch.cat(parts).float()


def _offset_rewards(offset, *, device="cpu"):
"""Groups with a large shared offset and unit-scale spread."""
spread = torch.tensor([0.0, 1.0, 2.0, 3.0])
groups = [offset + spread * scale for scale in (1.0, 0.5, 2.0)]
return torch.cat(groups).to(device)


_OFFSET_BOUNDS = [0, 4, 8, 12]


def _adv_tokens(batch):
sample_adv = _reference_group_advantages(batch.rewards, _SPP)
return (
Expand Down Expand Up @@ -127,6 +148,25 @@ def test_requires_exactly_one_group_spec():
op.group_advantages(rewards, samples_per_prompt=_SPP, group_boundaries=[0, 4, 8, 12])


@pytest.mark.parametrize("offset", [1e4, 1e5, 1e6])
def test_group_advantages_stable_under_large_reward_offset(offset):
op = NativeGRPOLossOp()
rewards = _offset_rewards(offset)
expected = _float64_group_advantages(rewards, _OFFSET_BOUNDS)
by_spp = op.group_advantages(rewards, samples_per_prompt=_SPP)
by_bounds = op.group_advantages(rewards, group_boundaries=_OFFSET_BOUNDS)
assert torch.allclose(by_spp, expected, atol=1e-3)
assert torch.allclose(by_bounds, expected, atol=1e-3)


def test_group_advantages_shift_invariant():
op = NativeGRPOLossOp()
rewards = _batch(seed=7).rewards
base = op.group_advantages(rewards, samples_per_prompt=_SPP)
shifted = op.group_advantages(rewards + 1e4, samples_per_prompt=_SPP)
assert torch.allclose(shifted, base, atol=1e-2)


# pure-PyTorch reference op (loss from logits)
def test_forward_loss_matches_reference():
op = NativeGRPOLossOp()
Expand Down Expand Up @@ -258,6 +298,16 @@ def test_triton_group_advantages_matches_native():
assert torch.allclose(got_b, exp_b, atol=1e-5)


@requires_triton_cuda
@pytest.mark.parametrize("offset", [1e4, 1e5, 1e6])
def test_triton_group_advantages_stable_under_large_reward_offset(offset):
fused = TritonGRPOLossOp()
rewards = _offset_rewards(offset, device="cuda")
expected = _float64_group_advantages(rewards, _OFFSET_BOUNDS)
got = fused.group_advantages(rewards, group_boundaries=_OFFSET_BOUNDS).cpu()
assert torch.allclose(got, expected, atol=1e-3)


@requires_triton_cuda
def test_triton_apply_with_per_sequence_advantages_matches_native():
native = NativeGRPOLossOp()
Expand Down