Repository navigation
[AMD] gfx950 assembly attention for EAGLE verify, draft extend and decode - #37465
Merged
HaiShaw merged 6 commits intoSep 8, 2026
Merged
Conversation
zijiecode
marked this pull request as ready for review
September 1, 2026 19:42
zijiecode
requested review from
BBuf,
DarkSharpness,
Fridge003,
HaiShaw,
HydraQYH,
Qiaolin-Yu,
celve,
hebiao064,
ispobock,
merrymercy and
yuan-luo
as code owners
September 1, 2026 19:42
zijiecode
force-pushed
the
perf/mtp-verify-attn-aiter-asm
branch
from
September 2, 2026 23:51
766c021 to
5b8c5ec
Compare
zijiecode
force-pushed
the
perf/mtp-verify-attn-aiter-asm
branch
from
September 3, 2026 00:14
5b8c5ec to
94b5a70
Compare
zijiecode
force-pushed
the
perf/mtp-verify-attn-aiter-asm
branch
from
September 3, 2026 03:55
94b5a70 to
d3034fd
Compare
…ree gfx950 assembly kernel Hand-written MI35x (gfx950) assembly attention for head_dim 256 with the fp8 page-16 NHD KV cache, shipped as .s under sglang.kernels and assembled at first use with ROCm clang. Serves the EAGLE/MTP verify step, the draft-model extend (ragged q_len 1..4, tail-aligned in-kernel) and q_len-1 decode in the aiter attention backend; GQA ratios 16 and 8. On by default on gfx95x for supported shapes, SGLANG_ASM_VERIFY_ATTN=0 keeps the Triton path.
zijiecode
force-pushed
the
perf/mtp-verify-attn-aiter-asm
branch
from
September 3, 2026 05:56
d3034fd to
3b7e94a
Compare
…; gate on gfx950 exactly A failed clang run or hipModuleLoad now disables the kernel for the process (one warning line) and the current call continues on the Triton path instead of raising. asm_kernel_available() also requires gcnArchName gfx950, matching the -mcpu the code object is built for.
Use the existing get_hip_version() helper (the same idiom as the aiter bpreshuffle gfx95 gates) so the kernel is only enabled on the ROCm versions it was validated on; older ROCm stays on the Triton path.
HaiShaw
approved these changes
Sep 8, 2026
HaiShaw
left a comment
Collaborator
There was a problem hiding this comment.
LGTM, 1st gfx asm kernel in-tree! ROCm specific only.
3 of 5 tasks
5 tasks
This was referenced Sep 10, 2026
4 of 5 tasks
zijiecode
added a commit
to zijiecode/sglang
that referenced
this pull request
Sep 12, 2026
…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.
This was referenced Sep 13, 2026
This was referenced Sep 15, 2026
This was referenced Sep 16, 2026
4 of 5 tasks
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
On MI355X, one EAGLE/MTP decode cycle of Qwen3.5-397B runs three attention calls through Triton split-K kernels: the target verify step (4 query tokens per sequence), the draft-model extend (q_len 1 to 4 per sequence) and q_len-1 decode. All three use head size 256, FP8 KV cache with page size 16, and 16 or 8 query heads per KV head. The Triton verify kernel streams KV at about 2.2 TB/s and the draft extend path is slower still; together they are most of the decode cycle at long contexts.
This PR serves all three with one hand-written gfx950 assembly kernel (
vattn_asm) that reads the paged FP8 cache in place at about 5 TB/s, 2.0-3.4x the Triton path. The kernel ships as.ssource undersglang.kernelsand is assembled at first use by the ROCm clang that every gfx950 image includes (validated on ROCm 7.2.0 and 7.2.4; the ROCm 7.0.0 assembler produces byte-identical code, but the path is gated to 7.2 or newer), so nothing changes at build time and no aiter release is needed.Modifications
sglang/kernels/ops/attention/vattn_asm_gfx950/: the attention kernel (vattn3_core.s), its split-KV reduce (vred.s) and a small loader that assembles them with ROCm clang at first use and launches them throughhipModuleLaunchKernel. The kernel handles q_len 1 to 4 per sequence in-kernel, so the draft paths need no padding or gather and stay cuda-graph capture safe.unified_attention_3d_mtp.py: route the verify call to the assembly kernel when the shape is supported; add two thin entry points for draft extend (ragged q_len) and q_len-1 decode that return False and leave the batch to the existing Triton path when they cannot serve it.aiter_backend.py: call those entry points from the verify, draft_extend and decode paths on gfx950 (head size 256, page size 16, FP8 KV, no sliding window, softcap or sinks). The verify gate also accepts 8 query heads per KV head (TP4).gcnArchNamematch) with ROCm 7.2 or newer (get_hip_version() >= (7, 2, 0)) when ROCm clang is present;SGLANG_ASM_VERIFY_ATTN=0keeps everything on the Triton path. If assembling or loading the kernel fails at first use, it is disabled for the process with one warning line and the call continues on the Triton path. Other shapes and GPUs are unchanged.Accuracy Tests
Verify path: outputs are bit-identical to the kernel build validated in the previous revision of this PR, so its GSM8K and GPQA results carry over. Draft extend and decode paths: outputs are bit-identical to the verify path on the same rows, and cuda-graph replay is bit-identical to eager. With a deliberately broken assembler, the first call returns without error, the kernel is disabled for the process, and the output is bitwise identical to the Triton path.
GSM8K, 1319 questions, TP4, real EAGLE decoding (3 steps, top-k 1, 4 draft tokens),
--attention-backend aiter --page-size 16 --kv-cache-dtype fp8_e4m3:unified_attention_3d_mtp(Triton)vattn_asm, verify onlyvattn_asm, verify + draft + decode (this branch)GPQA diamond (198 questions, repeat 8), verify path:
unified_attention_3d_mtp(Triton, from #36330)Benchmarking and Profiling
All numbers on MI355X (gfx950, ROCm 7.2), Qwen3.5-397B-A17B-MXFP4: head_dim 256, fp8 page-16 KV.
Kernel microbenchmarks
Verify shape (4 query tokens per sequence), single GPU, cold L2 (16 rotating KV replicas, randomly permuted page tables), kernel plus segment reduce, per layer:
unified_attention_3d_mtp+unified_attention_3d_mtp_reduce_segmentsvattn_asm+vred_asmDraft shapes, single GPU, per layer:
unified_attention(Triton)vattn_asmA streaming-read microbenchmark puts the device's practical read ceiling at 5.9-6.9 TB/s; ablating the compute phases of
vattn_asmchanges its time by less than 5%, i.e. the kernel is bound by K/V traffic.End-to-end serving, real agentic traces (InferenceX)
Same tree, two runs differing only in
SGLANG_ASM_VERIFY_ATTN(0 = all Triton, 1 = this PR: verify, draft_extend and decode on the asm kernel). The server is started by the unchanged MI355X InferenceX recipe (https://github.com/SemiAnalysisAI/InferenceX/blob/main/benchmarks/single_node/agentic/qwen3.5_fp4_mi355x_sglang_mtp.sh, as used by theqwen3.5-fp4-mi355x-sglang-agentic-mtpentries of https://github.com/SemiAnalysisAI/InferenceX/blob/main/configs/amd-master.yaml), and aiperf replays real agent sessions from the publiccc-traces-weka-062126-256kdataset over a 1200 s window (AIPERF_EXPERIMENTAL_FAST=1, one warmup request per lane).TP2, concurrency 32, KV offloading to DRAM with HiCache (the
kv-offloading: dramTP2 entry of the sweep;KV_OFFLOADING=dram KV_OFFLOAD_BACKEND=hicache). Both runs: 0 failed requests, 94% prefix cache hit.SGLANG_ASM_VERIFY_ATTN=0SGLANG_ASM_VERIFY_ATTN=1Matched-request check (the 1336 requests, identified by session and turn, that both runs completed): decode time per token 17.81 -> 15.21 ms (-14.6%), request latency -13.2%, TTFT -2.1%. This removes the effect of trace-replay phase drift from the per-request numbers.
TP4, concurrency 16, no KV offloading (the
kv-offloading: noneTP4 entry of the sweep). Both runs: 0 failed requests.SGLANG_ASM_VERIFY_ATTN=0SGLANG_ASM_VERIFY_ATTN=1Matched-request check (692 requests completed by both runs): decode time per token 6.02 -> 4.60 ms (-23.7%), request latency -21.9%, TTFT -7.2%.
Note on metrics: TPOT and total throughput are the stable comparison metrics for this benchmark; the interactivity mean is dominated by streaming chunk-coalescing outliers (it averages per-request 1/ITL), so interactivity is quoted at p50/p90.
Checklist
CI States
Latest PR Test (Base): ✅ Run #34170334564
Latest PR Test (Extra): ❌ Run #34170334440
Latest PR Test (AMD ROCm 7.2): ❌ Run #34170334605