Skip to content

[ASM] [HIP] [JIT] [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) - #5190

Closed
zijiecode wants to merge 2 commits into
ROCm:mainfrom
zijiecode:perf/mtp-verify-attn-asm-gfx950
Closed

zijiecode wants to merge 2 commits into
ROCm:mainfrom
zijiecode:perf/mtp-verify-attn-asm-gfx950

Conversation

@zijiecode

@zijiecode zijiecode commented Sep 1, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

This kernel was developed specifically for SGLang serving of Qwen3.5-397B on MI355X: the speculative-decoding verify step (EAGLE / MTP, 4 query tokens per sequence with a causal mask over the last 4 positions) over SGLang's NHD page-16 fp8 KV cache with head_dim 256. SGLang currently runs this step through its Triton split-K kernel unified_attention_3d_mtp: at production shapes (16 sequences of 35k context) it streams KV at about 2.2 TB/s (131.6 us per layer), the largest single decode-side kernel gap versus B200 in our MI355X study, and the TP4 configuration (GQA ratio 8) falls back to an even slower kernel (about 197 us per layer).

aiter's paged-attention assembly family (hsa/gfx950/pa/*) only covers head_dim 128 and expects the vLLM 5D KV layout, so nothing off-the-shelf applies. This PR adds mtp_verify_attn_fwd_asm: a hand-written gfx950 kernel for head_dim 256 that reads the NHD page-16 fp8 cache in place (no re-layout, no gather), supports GQA ratios 16 and 8 natively, and streams KV at about 5 TB/s, 2.0-3.4x the Triton path and about 80-85% of the streaming-read ceiling measured on the device. The matching SGLang integration (one call site) is sgl-project/sglang#37465.

Modifications

Follows the existing ASM op conventions (code objects under hsa/gfx950/<op>/, AiterAsmKernel launcher in csrc/py_itfs_cu/asm_*.cu, ctypes FFI op in aiter/ops/attention.py, op_tests/ test):

  • hsa/gfx950/mtp_verify_attn/: vattn_hd256_fp8_gqa16.co, vattn_hd256_fp8_gqa8.co (the main kernel vattn_asm, one code object per GQA ratio) and vattn_reduce.co (the split-KV segment reduce vred_asm).
  • csrc/py_itfs_cu/asm_mtp_verify_attn.cu: AITER_C_ITFS void mtp_verify_attn_fwd(...), argument validation, kernarg packing, AiterAsmKernel launches of the main kernel (grid = segments x sequences x kv heads) and the reduce.
  • aiter/ops/attention.py: _mtp_verify_attn_fwd ctypes stub, mtp_verify_attn_num_segments(), and the user-facing mtp_verify_attn_fwd_asm(q, k_cache, v_cache, block_tables, seq_lens, cu_seqlens_q, k_descale, v_descale, softmax_scale, num_segments=None, out=None) which allocates the split-KV workspaces and returns the bf16 output.
  • aiter/jit/optCompilerConfig.json: module_mtp_verify_attn_asm.
  • op_tests/test_mtp_verify_attn_asm.py: correctness against a fp32 torch reference (shuffled page tables, GQA 16 and 8, 2 KV heads, per-tensor descales != 1, ragged page counts with empty tail segments) and a --perf bandwidth sweep.

Kernel notes: one 8-wave work-group per (KV segment, sequence, KV head), one per CU; K/V tiles are DMA'd straight into LDS with global_load_lds_dwordx4, QK runs on v_mfma_f32_16x16x128_f8f6f4, PV reads its transposed V operand with ds_read_b64_tr_b8; Q and P are quantized to fp8 (the same precision class as the pa_* ASM kernels); K/V descales are read from device pointers so launches are HIP-graph-capture safe; empty segments are written neutrally so any segment count in [1, 64] is valid.

Supported: gfx950, head_dim 256, page size 16, K/V fp8 e4m3 with per-tensor scales, Q bf16, uniform q_len 4, GQA ratio 16 or 8. Everything else is rejected with AITER_CHECK.

Microbenchmark

MI355X (gfx950, ROCm 7.2), cold L2 (8-16 rotating KV replicas, randomly permuted page tables), main kernel plus reduce, per call. A verify batch of N sequences, each with context C and 4 query tokens; at TP2 each rank has 16 query heads on 1 KV head (GQA 16), at TP4 8 query heads on 1 KV head (GQA 8). Effective bandwidth is the KV bytes read (N x C x 256 x 2 bytes) divided by the kernel time.

Kernel 16 seqs x 35k ctx, GQA 16 16 x 70k, GQA 16 16 x 35k, GQA 8 4 x 35k, GQA 16
SGLang Triton unified_attention_3d_mtp + reduce (reference) 131.6 us (2.2 TB/s) 261 us (2.2 TB/s) 197 us (fallback kernel) -
mtp_verify_attn_fwd_asm (vattn_asm + vred_asm) 67 us (5.0 TB/s kernel-only) 105 us (5.3 TB/s) 57 us (5.5 TB/s) 30 us

Correctness against a fp32 torch reference (shuffled page tables, GQA 16 / GQA 8 / 2 KV heads, k/v descale != 1): max error 0.13 sigma of the output, the fp8 P/Q quantization floor. Ablating the compute phases of the kernel changes its time by less than 5%, i.e. it is bound by the K/V traffic; a streaming-read microbenchmark puts this device's practical read ceiling at 5.9-6.9 TB/s.

Accuracy in serving (SGLang, Qwen3.5-397B-A17B-MXFP4, TP4, real EAGLE decoding): GSM8K 5-shot on all 1319 questions is 1211/1319 with this kernel versus 1215/1319 with the Triton kernel.

Output of op_tests/test_mtp_verify_attn_asm.py --perf on MI355X:

Correctness (error relative to a fp32 torch reference, in units of the output's standard deviation):

ctx seqs q_heads kv_heads k_descale v_descale max_err_over_std
1000 2 16 1 1.0 1.0 0.075
1000 2 16 1 0.5 2.0 0.075
1000 2 8 1 1.0 1.0 0.075
777 2 32 2 1.0 1.0 0.083
8000 3 16 1 1.0 1.0 0.090

Bandwidth sweep (the full mtp_verify_attn_fwd_asm op including workspace allocation and the reduce, 8 rotating KV replicas, 64 iterations):

ctx seqs q_heads kv_heads segments us kv_TBps
35000 16 16 1 16 63.9 4.49
70000 16 16 1 16 108.2 5.30
35000 16 8 1 16 58.3 4.92
35000 4 16 1 64 30.6 2.35

Test Plan

op_tests/test_mtp_verify_attn_asm.py covers correctness against a fp32 torch reference (shuffled page tables, GQA 16 and 8, 2 KV heads, per-tensor descales, ragged page counts with empty tail segments) and, with --perf, the bandwidth sweep above:

python3 op_tests/test_mtp_verify_attn_asm.py --perf

End-to-end accuracy through SGLang (GSM8K and GPQA on Qwen3.5-397B, on par with the Triton kernel) is reported in the SGLang PR.

…V (mtp_verify_attn_fwd_asm)

Speculative-decoding verify attention (uniform q_len 4) over an NHD page-16
fp8 KV cache with head_dim 256, GQA ratio 16 or 8, developed for SGLang
serving of Qwen3.5 on MI355X. Split-KV main kernel (vattn_asm) plus segment
reduce (vred_asm), ~5 TB/s effective KV bandwidth, 2.0-3.4x SGLang's Triton
unified_attention_3d_mtp at production shapes. Follows the existing ASM op
conventions: code objects under hsa/gfx950/mtp_verify_attn/, AiterAsmKernel
launcher in csrc/py_itfs_cu/asm_mtp_verify_attn.cu, ctypes op
aiter.mtp_verify_attn_fwd_asm, op_tests/test_mtp_verify_attn_asm.py.
@github-actions

github-actions Bot commented Sep 1, 2026

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs; main branch always runs both MI35X and MI300X
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 5190 --add-label <label>

PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to opt this PR out.

@zijiecode zijiecode changed the title [gfx950] Assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) Sep 1, 2026
@zijiecode
zijiecode marked this pull request as ready for review September 1, 2026 19:41
@zijiecode
zijiecode requested a review from a team September 1, 2026 19:41
@github-actions github-actions Bot changed the title [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) [ASM] [HIP] [JIT] [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) Sep 1, 2026
@zufayu
zufayu requested a review from amd-ruitang3 September 2, 2026 01:21
@zijiecode zijiecode changed the title [ASM] [HIP] [JIT] [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) [ASM] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) Sep 2, 2026
@zijiecode zijiecode changed the title [ASM] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) Sep 2, 2026
@zijiecode zijiecode changed the title gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) [ASM] [HIP] [JIT] [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm) Sep 2, 2026
@zufayu

zufayu commented Sep 2, 2026

Copy link
Copy Markdown
Contributor

aiter PR #5190 — [ASM] [HIP] [JIT] [Feature] gfx950 assembly MTP-verify attention for head_dim 256 fp8 paged KV (mtp_verify_attn_fwd_asm)

Adds a hand-written gfx950 assembly kernel for the speculative-decoding verify step (4 query tokens/sequence over an NHD page-16 fp8 KV cache, head_dim 256, GQA 16/8) with a split-KV main kernel plus segment reduce, exposed as aiter.mtp_verify_attn_fwd_asm.

⚠️ NEEDS WORK

⚠️ The .co filenames and kernel symbols don't follow aiter's paged-attention naming convention, and three independent naming schemes are mixed in one op. Every existing gfx950 paged-attention code object uses a single pa-rooted name that runs through the whole chain unchanged — e.g. pa_bf16_pertokenFp8_gqa8_1tg_4w_mtp_msk1.co, where pa is the op root, bf16/pertokenFp8/gqa8/1tg_4w are the config, and mtp is a feature suffix (not the op identity). This PR scatters at least four unrelated names across the same op:

layer this PR convention
.cu file asm_mtp_verify_attn.cu asm_pa_*.cu
op name mtp_verify_attn_fwd_asm pa_*_asm
.co dir mtp_verify_attn/ pa_*/
.co file vattn_hd256_fp8_gqa16.co pa_*_hd256_*_gqa16_*.co
kernel symbol vattn_asm pa_* (or _ZN5aiter...pa_...E)
.s source vattn3_core.s / vred.s same root as the .co
reduce vred_asm / vattn_reduce.co / vred.s one name

The reduce alone has three names (vred_asm symbol, vattn_reduce.co file, vred.s source). The .co filename prefix vattn_ doesn't match its own directory mtp_verify_attn/, and neither matches the kernel symbol. In the existing family the .co prefix always equals the directory equals the op root (pa/pa/pa). Author should agree with @amd-ruitang3 (tangrui, top attention.py ASM committer) on a single pa-rooted name and apply it through the .cu op name, .co directory, .co filenames, and kernel symbols — e.g. along the lines of pa_bf16_pertokenFp8_hd256_page16_gqa16_8w_mtp.co / pa_bf16_pertokenFp8_hd256_page16_gqa8_8w_mtp.co for the main kernel and pa_reduce_hd256.co for the reduce, keeping mtp as a feature suffix the way pa_bf16_pertokenFp8_gqa8_1tg_4w_mtp_msk1.co already does. This is a rename + .co rebuild, not a behavior change.

@HaiShaw HaiShaw closed this Sep 3, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants