Repository navigation
[AMD] gfx950 assembly attention: length-aware split-KV for dynamic workload - #39172
Merged
Merged
Conversation
zijiecode
requested review from
BBuf,
DarkSharpness,
Fridge003,
HaiShaw,
HydraQYH,
Qiaolin-Yu,
celve,
hebiao064,
ispobock,
merrymercy and
yuan-luo
as code owners
September 12, 2026 04:02
…entic batches For bs > 1 a small Triton plan kernel picks the segment length from the batch's sequence lengths and writes a compact work list; the asm attention kernel runs as a 1-D grid over that list (idle work-groups exit before touching memory) and the reduce reads a per-token segment count. The plan is built once per forward and cached on the seq_lens / cu_seqlens_q tensors (keyed on the stream-capturing flag). bs == 1 keeps the fixed split of sgl-project#37465 (plan_ptr == 0), which already uses 64 segments there.
zijiecode
force-pushed
the
pr/vattn-segplan
branch
from
September 12, 2026 04:49
78badc3 to
e899634
Compare
This was referenced Sep 13, 2026
yichiche
approved these changes
Sep 14, 2026
yichiche
left a comment
Collaborator
There was a problem hiding this comment.
Split-K optimization for the MTP assembly kernel; changes are AMD-only. LGTM.
HaiShaw
approved these changes
Sep 14, 2026
5 tasks done
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
#37465 splits every sequence of a verify / draft batch into the same number of KV segments,
max(1, min(64, pow2floor(256 / (bs * kv_heads)))), one work-group per (segment, sequence, kv head). That is the right split when all sequences have similar length, not optimal when one batch holds requests that varies heavily in length. Two things then happen:This PR picks the segment length from the batch's actual sequence lengths and hands the kernel a compact work list, so every sequence gets as many segments as its length needs while the total work-group count stays at one per CU.
Modifications
The planned split is the default for bs > 1; bs == 1 keeps the fixed split of #37465 (
plan_ptr == 0), which already uses 64 segments there. No new switch. Kernel build and the gfx950 / ROCm gating are unchanged from #37465.vattn_asm_gfx950/__init__.py: a small Triton plan kernel (grid = num_seqs, graph-capture safe) picks one segment lengthTfromseq_lens(smallestTwith total segments <= num_CUs / kv_heads and per-sequence segments <=seg_max) and writes a work listplan[1 + i] = seq << 16 | seg(-1past the end) plus a per-query-token segment count. The attention kernel is launched as a 1-D grid over the work list;bs == 1bypasses the plan. The plan is built once per forward (cached on theseq_lens/cu_seqlens_qtensors, cleared at the start of every forward, keyed on the stream-capturing flag).vattn3_core.s: kernarg 128 -> 136 B (plan_ptr); the prologue decodes(seq, seg)from the work list and exits on-1before touching memory; tiles per segment come fromplan[0].vred.s: kernarg 56 -> 64 B (tok_nseg_ptr); the reduce reads only that token's segment count.unified_attention_3d_mtp.py/aiter_backend.py:reset_verify_attn_plan_cache()called frominit_forward_metadataandinit_forward_metadata_out_graph.Accuracy Tests
test/registered/kernel/attention/test_vattn_segplan.py(stage-b-test-1-gpu-small-amd-mi35x, 22 s on MI355X): 16 cases (GQA 16 and 8; uniform 16x70k; agentic skew including 0/1/5/17-token sequences; bs 1 / 2 / 24 / 64; ragged q 1..4) against an fp32 torch reference and against the fixed split. Same error vs. reference as the fixed split (max 0.0141 bf16 on the skew case, dominated by the fp8 KV); plan vs. fixed split differ by at most 0.007 (different fp32 reduction order across segment boundaries). Plan contents (T, work list,tok_nseg) checked against a Python reference in every case.seq_lenschange or a different tensor -> recompute; cached output bit-identical to uncached.Benchmarking and Profiling
unified_attention_3d_mtpEnd to end, InferenceX agentic recipe (Qwen3.5-397B-A17B MXFP4, MTP, aiperf trace replay; official harness mode: full aiperf, 3600 s, 10 warmups), MI355X, same tree and flags; off = fixed split of #37465, on = this PR. TP4 pairs ran one server at a time on the same GPUs; TP2 pairs ran off and on simultaneously on separate GPU pairs of the same node.
Server-side accept length is identical between off and on at every point (3.38-3.40). Same-tree A/B via a development-only switch that this PR does not ship. The p90 tail improves by 4-5% at every point (a long request joining a small batch no longer runs as 8 or fewer oversized segments); the median moves only at TP2 c32 (TPOT p50 12.47 -> 11.57 ms, -7.2%), where half of the decode steps are at bs 17-30 and the fixed split degrades. A same-configuration repeat on this replay differs by about 0.8%.
Checklist
test/registered/kernel/attention/test_vattn_segplan.py).CI States
Latest PR Test (Base): ✅ Run #34686493077
Latest PR Test (Extra): ❌ Run #34686492967
Latest PR Test (AMD ROCm 10): ❌ Run #34686493122