Repository navigation
Conversation
…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.
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
PR title tags & labels: |
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
The reduce alone has three names ( |
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 addsmtp_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>/,AiterAsmKernellauncher incsrc/py_itfs_cu/asm_*.cu, ctypes FFI op inaiter/ops/attention.py,op_tests/test):hsa/gfx950/mtp_verify_attn/:vattn_hd256_fp8_gqa16.co,vattn_hd256_fp8_gqa8.co(the main kernelvattn_asm, one code object per GQA ratio) andvattn_reduce.co(the split-KV segment reducevred_asm).csrc/py_itfs_cu/asm_mtp_verify_attn.cu:AITER_C_ITFS void mtp_verify_attn_fwd(...), argument validation, kernarg packing,AiterAsmKernellaunches of the main kernel (grid = segments x sequences x kv heads) and the reduce.aiter/ops/attention.py:_mtp_verify_attn_fwdctypes stub,mtp_verify_attn_num_segments(), and the user-facingmtp_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--perfbandwidth 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 onv_mfma_f32_16x16x128_f8f6f4, PV reads its transposed V operand withds_read_b64_tr_b8; Q and P are quantized to fp8 (the same precision class as thepa_*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.
unified_attention_3d_mtp+ reduce (reference)mtp_verify_attn_fwd_asm(vattn_asm+vred_asm)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 --perfon MI355X:Correctness (error relative to a fp32 torch reference, in units of the output's standard deviation):
Bandwidth sweep (the full
mtp_verify_attn_fwd_asmop including workspace allocation and the reduce, 8 rotating KV replicas, 64 iterations):Test Plan
op_tests/test_mtp_verify_attn_asm.pycovers 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: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.