Skip to content

fix(grpo-loss): compute group variance two-pass to avoid fp32 cancellation - #475

Closed
effintell wants to merge 4 commits into
RL-Align:mainfrom
effintell:fix/grpo-variance-cancellation
Closed

effintell wants to merge 4 commits into
RL-Align:mainfrom
effintell:fix/grpo-variance-cancellation

Conversation

@effintell

@effintell effintell commented Oct 5, 2026 •

Copy link
Copy Markdown

Summary

GRPO group advantages computed the per-group variance with the one-pass
form E[x^2] - E[x]^2. In fp32 this cancels catastrophically when the rewards
in a group share a large offset: with rewards 1e4 + [0, 1, 2, 3] the advantages
came out around ±5e5 instead of ±1.34, which silently blows up the policy loss.

This PR switches the native and Triton paths to a two-pass variance,
E[(x - mean)^2], matching what DistributedGRPOLossOp already does.
Advantages are now shift-invariant: adv(r + c) == adv(r) up to fp32 rounding of r + c.

Changes

  • Native (rl_engine/kernels/ops/pytorch/loss/grpo_loss.py): centre
    rewards by their group mean before squaring, and accumulate squared deviations with
    index_add_ instead of raw squared sums. Reuses the centred tensor for the final
    advantage.
  • Triton (rl_engine/kernels/ops/triton/loss/grpo_loss.py, _group_norm_kernel):
    same two-pass computation. Masked lanes are forced back to 0 after centring with
    tl.where(keep, rewards - mean, 0.0) so they don't add mean^2 to the variance.
  • Docs (docs/operators/grpo-loss.md): note the two-pass variance, the failure mode it fixes, and the shift-invariance guarantee.

Tests

New tests in tests/test_grpo_loss.py, all compared against a float64
two-pass reference (_float64_group_advantages):

  • test_group_advantages_stable_under_large_reward_offset[1e4|1e5|1e6]:
    native op, both the samples_per_prompt and group_boundaries paths, atol=1e-3.
  • test_group_advantages_shift_invariant: adv(r + 1e4) vs adv(r), atol=1e-2.
  • test_triton_group_advantages_stable_under_large_reward_offset[1e4|1e5|1 e6]: Triton kernel on CUDA (skipped without Triton + CUDA).

Test plan:

  • pytest tests/test_grpo_loss.py -v (CPU)
  • pytest tests/test_grpo_loss.py -v -k triton on a CUDA machine

Notes

  • Numerical results for well-conditioned rewards are unchanged up to fp32 rounding; existing native/Triton parity tests still apply.

Summary by CodeRabbit

  • Bug Fixes

    • Improved group-advantage normalization stability for rewards with large shared offsets across native, Triton, and distributed backends.
    • Preserved shift invariance up to floating-point rounding.
  • Tests

    • Added coverage for large reward offsets and comparisons against a higher-precision reference across supported normalization paths.
  • Documentation

    • Clarified the variance calculation and its numerical stability characteristics.

…ation

The one-pass E[x^2] - E[x]^2 form cancels catastrophically in fp32 when
rewards share a large offset. With rewards 1e4 + [0, 1, 2, 3] the
advantage came out around +/-5e5 instead of +/-1.34.

Centre the group before squaring in both the native and Triton paths. In
the Triton kernel, keep masked lanes at zero after centring with
tl.where(keep, ...).
Compare against a float64 two-pass reference, parameterised over reward
offsets of 1e4, 1e5 and 1e6. Also check shift-invariance, and add a
Triton CUDA variant.
@coderabbitai

coderabbitai Bot commented Oct 5, 2026 •

Copy link
Copy Markdown

Review in Change Stack →

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: defaults
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: b091a054-b0c1-492f-b501-b4c0382d5f21
📥 Commits

Reviewing files that changed from the base of the PR and between 43f150f and 14a518d.

📒 Files selected for processing (4)
  • docs/operators/grpo-loss.md
  • rl_engine/kernels/ops/pytorch/loss/grpo_loss.py
  • rl_engine/kernels/ops/triton/loss/grpo_loss.py
  • tests/test_grpo_loss.py

Included review availability: This review used your included allowance. Your plan provides up to 2 included reviews per hour; 1 remain after this review.


📝 Walkthrough

Walkthrough

The native and Triton GRPO loss kernels now calculate group variance from centered rewards. Tests check normalization against a float64 reference at large reward offsets and check shift invariance. The operator documentation describes the two-pass calculation.

Changes

GRPO group normalization

Layer / File(s) Summary
Two-pass variance and numerical validation
rl_engine/kernels/ops/pytorch/loss/grpo_loss.py, rl_engine/kernels/ops/triton/loss/grpo_loss.py, tests/test_grpo_loss.py, docs/operators/grpo-loss.md
Both kernels calculate population variance from centered rewards. Tests compare native and Triton results with a float64 reference for large offsets, and check native shift invariance. The documentation describes the two-pass calculation and its numerical behavior.

Priority: ➖ Normal

Estimated code review effort: 2 (Simple) | ~10 minutes

Change: Bug fix

Merge Risk: ⚪ Minimal · up to 74521

No outstanding issue identified in the GRPO variance change; it is ready to merge after normal checks.

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 35.71% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 14 functions across 3 files. (1 skipped: … Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly identifies the two-pass group variance change and its purpose: avoiding fp32 cancellation in GRPO loss.
Full details: Docstring Coverage

Explanation

Docstring coverage is 35.71% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 14 functions across 3 files. (1 skipped: 1 unsupported.)

  • Fix all pre-merge checks with AI
✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create a new PR
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Comment @coderabbitai help to get the list of available commands.

@Flink-ddd Flink-ddd added the bug Something isn't working label Oct 5, 2026
@effintell

Copy link
Copy Markdown
Author

I just noticed that this issue is already being addressed in #455.
I'll close this PR to avoid duplicating the work. Thanks!

@effintell effintell closed this Oct 6, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

bug Something isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants