Skip to content

[KDA] Zero every CUDA-graph padded request's output rows in the CuTe DSL MTP kernel - #41733

Open
rodamani wants to merge 3 commits into
sgl-project:mainfrom
modal-projects:rohan/up/kda-mtp-padded-row-zeroing
Open

rodamani wants to merge 3 commits into
sgl-project:mainfrom
modal-projects:rohan/up/kda-mtp-padded-row-zeroing

Conversation

@rodamani

@rodamani rodamani commented Sep 29, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

In the CuTe DSL KDA MTP verify kernel (kernels/ops/attention/kda_decode_mtp.py), graph-padded requests (ssm_state_indices == -1) all share the last real cu_seqlens endpoint, because the hybrid linear-attention backend collapses every padded request's query_start_loc to the real-token endpoint when num_padding > 0. The pad branch computed its output row block from cu_seqlens[i_n], so it cleared only the first padded request's rows and left the other padded requests' dense output rows as uninitialised torch.empty bytes.

We measured this in an instrumented engine run (no fix applied): the first padded request's output interval was zero in 100% of 5.3M observations across 8 ranks, while padded requests 2 and 3 were nonzero in 6.1%-59.0% of theirs.

Modifications

  • The wrapper enforces a fixed dense 1 + num_spec width per request, so the pad branch now clears row block i_n * T_LOOP.
  • test_cutedsl_multiple_padding_outputs_and_states in test/registered/kernels/ops/attention/test_kda_mtp_cutedsl_replayssm_ring.py: sentinel-filled padded outputs and state slots through the real wrapper, eager and CUDA-graph replay, snapshot and ReplaySSM-ring arms, raw and fused-onorm outputs, 1/2/3 padded requests. Real-request outputs must match a real-requests-only control run bitwise, padded outputs must be zero, and untouched state must be unchanged.

Overlap: #41714 also edits this kernel and test file (BF16 recurrent state). This PR uses the FP32 state that main supports today; whichever lands second needs a small rebase of the test helper.

Accuracy Tests

The new test needs SM100 and has not been run on this branch yet (it skips on our CPU environment). It passed on our internal branch, where the npad >= 2 cases fail without the kernel change. A Blackwell run on this branch is pending.

Speed Tests and Profiling

No hot-path change beyond the fix itself; not separately benchmarked.

Checklist

Review and Merge Process

  1. Ping Merge Oncalls to start the process. See the PR Merge Process.
  2. Get approvals from CODEOWNERS and other reviewers.
  3. Trigger CI tests with comments or contact authorized users to do so.
    • Common commands include /tag-and-rerun-ci, /tag-run-ci-label, /rerun-failed-ci
  4. After green CI and required approvals, ask Merge Oncalls or people with Write permission to merge the PR.

CI States

Latest PR Test (Base): ✅ Run #36767819535
Latest PR Test (Extra): ❌ Run #36767818987
Latest PR Test (AMD ROCm 10): ❌ Run #36767819493

…he CuTe DSL MTP kernel (sgl-project#138)

* kda: zero each padded CUDA-graph row's own dense output interval in the CuTe DSL MTP kernel

Graph-padded requests (ssm_state_indices == -1) all share the last real
cu_seqlens endpoint: production graph metadata collapses every padded
request's query_start_loc to the real-token endpoint
(hybrid_linear_attn_backend.py, verify branch with num_padding > 0), so the
kernel's pad branch cleared only the first padded request's T_LOOP rows and
left the remaining padded rows as torch.empty bytes. The wrapper enforces a
fixed dense 1 + num_spec width per request, so the pad branch now clears row
block i_n * T_LOOP.

The defect geometry is measured, not just inferred: an instrumented engine
run observed the first padded request's dense output interval zero in 100%
of 5.3M observations across 8 ranks while padded requests 2 and 3 were
nonzero in 6.1%-59.0% of theirs, with zero nonfinite values in every arm.
That experiment applied no fix.

Adds test_cutedsl_multiple_padding_outputs_and_states: sentinel-filled
padded output intervals and state slots exercised through the real wrapper,
in eager and CUDA-graph replay, with both snapshot and ReplaySSM-ring state
arms and both raw and fused-onorm outputs; real-request outputs must match a
real-requests-only control run bitwise and all untouched padding state must
stay unchanged.

Co-Authored-By: Rahul Chalamala <22563365+rchalamala@users.noreply.github.com>

* test: cover one, two, and three padded requests in the KDA CuTeDSL multiple-padding test

Extends test_cutedsl_multiple_padding_outputs_and_states to the three
controlled conditions of that measurement: 7/6/5 real of 8 padded rows on
a block-8 shape (npad 1/2/3), keeping the original 2-of-4 small-shape case.
The npad-1 arm guards the already-correct single-padding geometry against
regression; the npad>=2 arms fail on the unpatched kernel, which is what the
measurement observed (first padded interval always zero, later padded
intervals carrying torch.empty bytes).

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
@rodamani
rodamani marked this pull request as ready for review September 29, 2026 21:16
@rodamani

Copy link
Copy Markdown
Contributor Author

/rerun-test -c test_kda_decode_mtp.py

@github-actions

github-actions Bot commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

Results for /rerun-test -c test_kda_decode_mtp.py:

🚀 4-gpu-b200 (2 tests): ✅ View workflow run

cd test/ && python3 registered/kernels/ops/attention/test_kda_decode_mtp.py
cd test/ && python3 registered/kernels/ops/attention/test_kda_mtp_cutedsl_replayssm_ring.py

@rodamani

Copy link
Copy Markdown
Contributor Author

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Sep 30, 2026

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

jit-kernel run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants