Skip to content

fix(ratio_kl): isolate inactive logits before normalization - #465

Open
MichaelCao0 wants to merge 3 commits into
RL-Align:mainfrom
MichaelCao0:codex/fix-ratio-kl-masked-grad
Open

MichaelCao0 wants to merge 3 commits into
RL-Align:mainfrom
MichaelCao0:codex/fix-ratio-kl-masked-grad

Conversation

@MichaelCao0

@MichaelCao0 MichaelCao0 commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

NativeRatioKLOp previously 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:

  • Original implementation: 8 failing CPU regressions and 24 failing GPU regressions, all exposing inactive NaN gradients.
  • Fixed CPU subset: 24 passed; complete RatioKL suite: 150 passed, including actual Native/Triton FP32, FP16 and BF16 checks.
  • Changed-file Black, isort, Ruff, Flake8 and whitespace checks passed.
  • Strict documentation builds on both base and head stop on the same eight pre-existing broken links; head adds no warnings.

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

  • Bug Fixes
    • Inactive actions now produce a ratio of 1, a KL penalty of 0, and zero policy-logit gradients, even when their logits or action IDs are invalid.
    • Invalid logits on active actions remain visible in results and gradients; out-of-range active action IDs are rejected.
  • Documentation
    • Clarified inactive-action behavior, differences in how native and Triton backends handle inactive rows, and how the native reference handles empty or all-inactive inputs and normalizes rows.

Signed-off-by: MichaelCaoo <139663530+MichaelCao0@users.noreply.github.com>
@coderabbitai

coderabbitai Bot commented Oct 3, 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: e1088870-486e-41d6-97bc-d8c66d932f0a
📥 Commits

Reviewing files that changed from the base of the PR and between 7fbbd91 and b7e218c.

📒 Files selected for processing (1)
  • rl_engine/kernels/ops/pytorch/loss/ratio_kl.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 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.

Changes

Ratio-KL inactive-row handling

Layer / File(s) Summary
Mask inactive logits before normalization
rl_engine/kernels/ops/pytorch/loss/ratio_kl.py, docs/operators/ratio-kl.md
The native operator masks inactive policy and reference logits before normalization. Documentation specifies neutral inactive outputs, zero policy-logit gradients, active-row behavior, and native reference performance details.
Verify inactive and active rows
tests/test_ratio_kl.py
Tests compare active outputs and gradients with an active-row reference. They cover inactive logits and action IDs, empty and all-masked batches, non-contiguous inputs, shared-weight gradients, non-finite active logits, and out-of-range active action IDs across supported native and Triton paths.

Priority: ➖ Normal

Estimated code review effort: 3 (Moderate) | ~20 minutes

Change: Bug fix

Merge Risk: ⚪ Minimal · up to b7e21

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)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 18.18% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 11 functions across 2 files. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly and concisely describes the main change: masking inactive logits before normalization in the ratio/KL operator.
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.
  • 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 4, 2026

@Flink-ddd Flink-ddd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

@MichaelCao0

Copy link
Copy Markdown
Contributor Author

I reran the before/after regressions and added a diagnostic that reports the affected policy-logits and shared-weight gradients. The attached pr-465-evidence-v2.zip includes full logs, exact commands and exit codes, source hashes, and reproduction scripts.

The revisions used were:

  • Baseline: 11cac8c46fa9f6fae67ffd3b4054c11337d63c04
  • Original fix: b84f19b6ac89aecd242ff191ca928819ed768b57
  • Current PR head: 7fbbd91c5465fd99a65dff9f1610eb34ba719917

For the baseline run, I copied only tests/test_ratio_kl.py from the PR head onto the baseline source. Both runs use identical regression tests.
From each source directory, with PYTHONPATH="$PWD" and PYTEST_DISABLE_PLUGIN_AUTOLOAD=1:

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
Check Before After
Native CPU regressions 8 failed, 5 passed 13 passed
CUDA regressions 24 failed, 54 passed 78 passed
Complete RatioKL file — 150 passed

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:

test_inactive_logits_do_not_contaminate_gradients[native-cpu-fp32-False-nan]
    assert torch.isfinite(policy.grad).all()
E   assert tensor(False)

test_inactive_logits_do_not_contaminate_shared_weight_gradient[native-cpu-fp32--inf]
    assert torch.isfinite(weight.grad).all()
E   assert tensor(False)

The attached diagnostic uses finite hidden[6,4] and trainable weight[4,17], computes policy = hidden @ weight + padding, and adds either NaN or -inf to three inactive rows. An independent active-only reference computes selected log-probabilities using selected_logit - logsumexp(logits) and backpropagates into a separate copy of the weights.

For all four FP32 combinations—CPU/CUDA × inactive NaN/-inf—the results were:

Diagnostic Before After
Inactive forward values ratio=1, KL=0 ratio=1, KL=0
Non-finite policy-logits gradient elements 51 0
Non-finite shared-weight gradient elements 68 / 68 0 / 68
Inactive policy-logits gradients exactly zero false true
Shared-weight gradients match active-only reference false true

After the fix, the maximum absolute shared-weight gradient error against that reference was 2.086162567138672e-07. Reference logits remained detached.

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.sh

The wrapper exports the pinned revisions into temporary directories and runs the regressions plus probe_ratio_gradients.py. I also executed the exact wrapper extracted from the original ZIP in a separate checkout; it reproduced the same results. The additional session is included under attachment-replay/.

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.
Uploading pr-465-evidence-v2.zip…

@Flink-ddd
Flink-ddd requested a review from ryankert01 as a code owner October 8, 2026 01:36

@Flink-ddd Flink-ddd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks, the before and after results address my review. LGTM.

This branch has not been deployed

No deployments
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