Repository navigation
Conversation
…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
marked this pull request as ready for review
September 29, 2026 21:16
rodamani
requested review from
BBuf,
DarkSharpness,
HaiShaw,
HydraQYH,
celve and
yuan-luo
as code owners
September 29, 2026 21:16
Contributor
Author
|
/rerun-test -c test_kda_decode_mtp.py |
Contributor
|
Results for 🚀 |
Contributor
Author
|
/tag-and-rerun-ci |
3 of 5 tasks
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 realcu_seqlensendpoint, because the hybrid linear-attention backend collapses every padded request'squery_start_locto the real-token endpoint whennum_padding > 0. The pad branch computed its output row block fromcu_seqlens[i_n], so it cleared only the first padded request's rows and left the other padded requests' dense output rows as uninitialisedtorch.emptybytes.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
1 + num_specwidth per request, so the pad branch now clears row blocki_n * T_LOOP.test_cutedsl_multiple_padding_outputs_and_statesintest/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
/tag-and-rerun-ci,/tag-run-ci-label,/rerun-failed-ciCI States
Latest PR Test (Base): ✅ Run #36767819535
Latest PR Test (Extra): ❌ Run #36767818987
Latest PR Test (AMD ROCm 10): ❌ Run #36767819493