Repository navigation
fix(ratio_kl): isolate inactive logits before normalization - #465
MichaelCao0 wants to merge 3 commits into
Conversation
Signed-off-by: MichaelCaoo <139663530+MichaelCao0@users.noreply.github.com>
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (1)
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 ratio-KL operator masks inactive policy and reference logits before normalization. Documentation and tests cover inactive-row outputs and gradients, invalid inactive action IDs, and validation of active action IDs. ChangesRatio-KL inactive-row handling
Priority: ➖ Normal Estimated code review effort: 3 (Moderate) | ~20 minutes Change: Bug fix Merge Risk: ⚪ Minimal · up to The documented inactive-row behavior and active-row checks are covered by the changed tests, and the RatioKL suite is reported passing. No specific merge blocker remains. 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
Flink-ddd
left a comment
There was a problem hiding this comment.
Could you share the test commands and before and after logs showing that inactive NaN -inf rows no longer contaminate policy or shared-weight gradients?
|
I reran the before/after regressions and added a diagnostic that reports the affected policy-logits and shared-weight gradients. The attached The revisions used were:
For the baseline run, I copied only python -m pytest tests/test_ratio_kl.py \
-q --tb=short -p no:cacheprovider \
-k 'inactive_logits and native-cpu'
python -m pytest tests/test_ratio_kl.py \
-q --tb=short -p no:cacheprovider \
-k 'inactive_logits and cuda'
# Complete test file, fixed version.
python -m pytest tests/test_ratio_kl.py \
-q --tb=short -p no:cacheprovider
The CUDA selection includes Native and Triton in FP32, FP16, and BF16. All 24 pre-fix CUDA failures are in Native; Triton already passes these inactive-row cases. The complete-file total includes the focused tests above. Representative baseline failures are: The attached diagnostic uses finite For all four FP32 combinations—CPU/CUDA × inactive NaN/
After the fix, the maximum absolute shared-weight gradient error against that reference was This demonstrates the backward issue directly: neutral masked forward outputs were already present before the fix, but non-finite log-softmax intermediates still contaminated gradients. Masking inactive logits before normalization removes that contamination. Separate tests verify that invalid active distributions remain visible. To reproduce, extract the attachment and run this from an RL-Kernel Git clone with a prepared Python environment and a free GPU: RLK_PY=python CUDA_VISIBLE_DEVICES=0 \
bash /path/to/extracted/reproduce-pr-465.shThe wrapper exports the pinned revisions into temporary directories and runs the regressions plus The shared-weight result applies to the tested finite upstream matmul inputs with non-finite additive padding afterward. It does not imply that zero gradients at this operator's boundary can repair arbitrary non-finite computations earlier in a model. |
Flink-ddd
left a comment
There was a problem hiding this comment.
Thanks, the before and after results address my review. LGTM.
NativeRatioKLOppreviously evaluated log-softmax on inactive logits before masking the selected log-prob differences. An inactive row containing NaN or only negative infinity produced neutral forward values but NaN policy-logits gradients, which could propagate to shared parameter gradients.Mask inactive policy and reference logits before normalization. Active values and arithmetic, reference freezing, and action-ID validation remain intact. Add independent active-row reference tests for forward/backward behavior, shared-weight gradients, strided inputs, empty/all-inactive inputs, and active invalid distributions; document the inactive-row contract and allocation tradeoff.
Validation on NVIDIA H200 with PyTorch 2.8.0+cu128 and Triton 3.4.0:
The PyTorch reference uses full-size temporary masked copies and still normalizes all rows; no latency or memory improvement is claimed. Triton is unchanged. Full training and ROCm were not run.
Summary by CodeRabbit