Skip to content

[AMD] gfx950 assembly attention for EAGLE verify, draft extend and decode - #37465

Merged
HaiShaw merged 6 commits into
sgl-project:mainfrom
zijiecode:perf/mtp-verify-attn-aiter-asm
Sep 8, 2026
Merged

HaiShaw merged 6 commits into
sgl-project:mainfrom
zijiecode:perf/mtp-verify-attn-aiter-asm

Conversation

@zijiecode

@zijiecode zijiecode commented Sep 1, 2026 •

Copy link
Copy Markdown
Collaborator

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 .s source under sglang.kernels and 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

  • Add 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 through hipModuleLaunchKernel. 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).
  • On by default on gfx950 (exact gcnArchName match) with ROCm 7.2 or newer (get_hip_version() >= (7, 2, 0)) when ROCm clang is present; SGLANG_ASM_VERIFY_ATTN=0 keeps 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:

Attention kernel Accuracy Invalid
unified_attention_3d_mtp (Triton) 0.921 0.006
vattn_asm, verify only 0.918 0.005
vattn_asm, verify + draft + decode (this branch) 0.919 0.005

GPQA diamond (198 questions, repeat 8), verify path:

Validation Result
Mean score 0.879
unified_attention_3d_mtp (Triton, from #36330) 0.879
Reference (Qwen/Qwen3.5-397B-A17B model card) 0.884

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:

Kernel 16 x 35k, GQA 16 16 x 70k, GQA 16 16 x 35k, GQA 8
unified_attention_3d_mtp + unified_attention_3d_mtp_reduce_segments 131.6 us (2.2 TB/s) 261 us (2.2 TB/s) 197 us (fallback)
vattn_asm + vred_asm 67 us (5.0 TB/s) 105 us (5.3 TB/s) 57 us (5.5 TB/s)

Draft shapes, single GPU, per layer:

Batch unified_attention (Triton) vattn_asm
draft_extend, 8 seqs x 105k, ragged q_len 1..4 396 us 226 us
decode q_len 1, 8 seqs, contexts 6 to 105k 100.5 us 66.2 us

A streaming-read microbenchmark puts the device's practical read ceiling at 5.9-6.9 TB/s; ablating the compute phases of vattn_asm changes 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 the qwen3.5-fp4-mi355x-sglang-agentic-mtp entries of https://github.com/SemiAnalysisAI/InferenceX/blob/main/configs/amd-master.yaml), and aiperf replays real agent sessions from the public cc-traces-weka-062126-256k dataset over a 1200 s window (AIPERF_EXPERIMENTAL_FAST=1, one warmup request per lane).

aiperf profile --scenario inferencex-agentx-mvp \
    --url http://localhost:$PORT --endpoint /v1/chat/completions --endpoint-type chat --streaming \
    --model amd/Qwen3.5-397B-A17B-MXFP4 --tokenizer amd/Qwen3.5-397B-A17B-MXFP4 --tokenizer-trust-remote-code \
    --public-dataset semianalysis_cc_traces_weka_062126_256k --num-dataset-entries 393 --apply-chat-template \
    --concurrency $CONC --benchmark-duration 1200 --warmup-requests-per-lane 1 --warmup-grace-period 1800 \
    --trajectory-start-min-ratio 0.25 --trajectory-start-max-ratio 0.75 --trace-idle-gap-cap-seconds 300 \
    --random-seed 42 --failed-request-threshold 0.10 --use-server-token-count \
    --server-metrics http://localhost:$PORT/metrics --stats-interval 30 --slice-duration 1.0 --no-gpu-telemetry

TP2, concurrency 32, KV offloading to DRAM with HiCache (the kv-offloading: dram TP2 entry of the sweep; KV_OFFLOADING=dram KV_OFFLOAD_BACKEND=hicache). Both runs: 0 failed requests, 94% prefix cache hit.

TP2 c32 (HiCache) SGLANG_ASM_VERIFY_ATTN=0 SGLANG_ASM_VERIFY_ATTN=1 Delta
TPOT p50 16.34 ms 13.99 ms -14.4%
TPOT mean 20.93 ms 17.83 ms -14.8%
TPOT p90 25.70 ms 22.71 ms -11.7%
Interactivity p50 (tok/s per user) 61.2 71.5 +16.9%
Interactivity p90 95.5 120.7 +26.5%
Output throughput 828 tok/s 894 tok/s +8.0%
Requests completed in the window 1336 1433 +7.3%
E2E latency p90 34.2 s 28.0 s -18.0%
TTFT p50 1.33 s 1.27 s -5.0%

Matched-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: none TP4 entry of the sweep). Both runs: 0 failed requests.

TP4 c16 SGLANG_ASM_VERIFY_ATTN=0 SGLANG_ASM_VERIFY_ATTN=1 Delta
TPOT p50 6.15 ms 4.58 ms -25.6%
TPOT mean 7.47 ms 6.39 ms -14.4%
TPOT p90 8.48 ms 6.53 ms -23.0%
Interactivity p50 (tok/s per user) 162.7 218.6 +34.4%
Interactivity p90 249.5 298.4 +19.6%
Output throughput 550 tok/s 586 tok/s +6.5%
Requests completed in the window 692 736 +6.4%
E2E latency p90 14.3 s 10.3 s -28.4%
TTFT p50 0.68 s 0.61 s -9.8%

Matched-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

@zijiecode zijiecode changed the title [AMD] Use aiter's gfx950 assembly MTP-verify attention kernel (head_dim 256, fp8 paged KV) in the aiter backend [AMD][gfx950] Assembly MTP-verify attention kernel for Qwen3.5 Sep 1, 2026
@zijiecode
zijiecode marked this pull request as ready for review September 1, 2026 19:42
@yichiche yichiche added run-ci CI: run the baseline test suite on this PR amd-aiter-not-ready labels Sep 2, 2026
@zijiecode
zijiecode force-pushed the perf/mtp-verify-attn-aiter-asm branch from 766c021 to 5b8c5ec Compare September 2, 2026 23:51
@zijiecode zijiecode changed the title [AMD][gfx950] Assembly MTP-verify attention kernel for Qwen3.5 [AMD] gfx950 assembly attention for EAGLE verify, draft extend and decode (head_dim 256, fp8 paged KV) Sep 2, 2026
@zijiecode
zijiecode force-pushed the perf/mtp-verify-attn-aiter-asm branch from 5b8c5ec to 94b5a70 Compare September 3, 2026 00:14
@zijiecode zijiecode changed the title [AMD] gfx950 assembly attention for EAGLE verify, draft extend and decode (head_dim 256, fp8 paged KV) [AMD] gfx950 assembly attention for EAGLE verify, draft extend and decode Sep 3, 2026
@zijiecode
zijiecode force-pushed the perf/mtp-verify-attn-aiter-asm branch from 94b5a70 to d3034fd Compare September 3, 2026 03:55
…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
zijiecode force-pushed the perf/mtp-verify-attn-aiter-asm branch from d3034fd to 3b7e94a Compare September 3, 2026 05:56
zijiecode and others added 3 commits September 5, 2026 22:55
…; 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 HaiShaw 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.

LGTM, 1st gfx asm kernel in-tree! ROCm specific only.

@HaiShaw
HaiShaw merged commit e634ba7 into sgl-project:main Sep 8, 2026
313 of 356 checks passed
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.
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.

3 participants