Repository navigation
Conversation
…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.
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (4)
Included review availability: This review used your included allowance. Your plan provides up to 2 included reviews per hour; 1 remain after this review. 📝 WalkthroughWalkthroughThe 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. ChangesGRPO group normalization
Priority: ➖ Normal Estimated code review effort: 2 (Simple) | ~10 minutes Change: Bug fix Merge Risk: ⚪ Minimal · up to 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)
✅ Passed checks (4 passed)
Full details: Docstring CoverageExplanation 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.)
✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
|
I just noticed that this issue is already being addressed in #455. |
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 rewardsin a group share a large offset: with rewards
1e4 + [0, 1, 2, 3]the advantagescame out around
±5e5instead 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 whatDistributedGRPOLossOpalready does.Advantages are now shift-invariant:
adv(r + c) == adv(r)up to fp32 rounding ofr + c.Changes
rl_engine/kernels/ops/pytorch/loss/grpo_loss.py): centrerewards by their group mean before squaring, and accumulate squared deviations with
index_add_instead of raw squared sums. Reuses the centred tensor for the finaladvantage.
rl_engine/kernels/ops/triton/loss/grpo_loss.py,_group_norm_kernel):same two-pass computation. Masked lanes are forced back to
0after centring withtl.where(keep, rewards - mean, 0.0)so they don't addmean^2to the variance.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 float64two-pass reference (
_float64_group_advantages):test_group_advantages_stable_under_large_reward_offset[1e4|1e5|1e6]:native op, both the
samples_per_promptandgroup_boundariespaths,atol=1e-3.test_group_advantages_shift_invariant:adv(r + 1e4)vsadv(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 tritonon a CUDA machineNotes
Summary by CodeRabbit
Bug Fixes
Tests
Documentation