Pull requests / #329
#329 hip: gfx12 (RDNA4) WMMA kernel for the int8-KV QSA prompt attention (1.26-1.35x prompt processing)
closed · @bsorensen110 · 0 commentaires · Sur GitHub
AMD / HIPNVIDIA / CUDAModels & quants
Description
## What The int8-KV QSA prompt attention (`qsa_prompt_attn_batch`) has a tensor-core kernel on CUDA (sm_80 `mma.sync` + `cp.async`) that is compiled out on HIP, so AMD runs the ordered FP32 kernel. This adds the gfx12 (RDNA4: gfx1200/gfx1201) equivalent on `v_wmma_f32_16x16x16_f16`. Prompt processing on a Radeon AI PRO R9700 with Swift 1.5 IQ3_XXS, `--kv int8`: **1.35x at 31K tokens, 1.26x at 108K**. Decode is unchanged, because the change only touches the prompt path. ## How Same algorithm and accuracy contract as the CUDA kernel: FP16 hi+lo split of q and p, fixed summation order across the four 64-dim groups, online softmax, V scale folded into p. The parts that differ: - **Fragments.** gfx12 wave32 WMMA, verified against a host reference. A lane `l` holds row `l%16`, `k = 8*(l/16)+i`. B lane `l` holds column `l%16`, the same `k`. D lane `l` holds column `l%16`, rows `8*(l/16)+i`. The 12 heads fill A rows 0-11, rows 12-15 are zero. - **No `cp.async` on gfx12.** The next chunk's 16-byte loads stay in registers while the current chunk computes, and are written to the warp's own LDS slice at the next step. - **int8 to half.** Exact, via `v_perm_b32` + xor `0x64806480` + one packed subtract of 1152. Checked over all 65,536 byte pairs. - **Gating.** Compile time (`__gfx1200__` / `__gfx1201__`) and run time (`gcnArchName`). Every other HIP target, fp16 and K8V4 pools, and `STRATA_PROMPT_ATTN_OLD=1` keep the previous kernel. The CUDA path is unchanged. - `qsa_prompt_attn_parity` now builds on HIP, prints SKIP where the kernel is unavailable, and is registered with ctest. ## Measurements R9700 (gfx1201), ROCm nightly, Ryzen 9 9950X, default `--resident-cpu-experts`. Both arms ran from one binary with only the attention kernel switched (`STRATA_PROMPT_ATTN_OLD=1` for "old"). Synthetic parity harness, int8, 32K ctx x 2048 queries: | | old kernel | WMMA | |---|---|---| | time per chunk | 38.5 ms | 5.6 ms (**6.9x**) | | error vs FP64 (output scale 3.6) | 2.1e-6 | 2.6e-6 | | new vs old | | 1.7e-6 of scale | End to end (engine's own prefill timing, `STRATA_PREFILL_TIMING=1`): | prompt | old | WMMA | speedup | |---|---|---|---| | 30.8K tokens (3 runs) | 22.8 s | 16.9 s | 1.35x | | 107.6K tokens (1 run) | 100.7 s | 79.7 s | 1.26x | | 8.8K tokens (2 runs) | 7.1 s | 5.6 s | 1.28x | | 4.2K tokens (2 runs) | 3.7 s | 3.0 s | 1.24x | The QSA attention phase fell 6.5x (6.9 s to 1.1 s at 31K, 24.7 s to 3.8 s at 108K). Every other phase stayed within about 20 ms. At 108K the QSA select phase (not touched here) is now the largest remaining QSA cost, about 29% of the prompt. ## Testing - `ctest`: 43/43 pass on this tree (built from `v0.1.30`, no other patches), with the four standard exclusions (`ple_parity`, `platform_memory_test`, `expert_parity`, `pool_test`). `hip_prefill_hipblaslt_gemm` skips because no gfx1201 table ships. `qsa_prompt_attn_parity` is among the 43. - Needle recall on the real model: 4/4 with either kernel (32K at depths 10/50/90, one 124K prompt), identical answers. ## Limits - Replies are not bitwise identical between the kernels (as with the CUDA kernel, `qsa_prompt_attn.hpp`). I have not run a teacher-forced top-1 comparison or an old-vs-old control, so the needle result plus the parity harness are the correctness evidence here. - Only int8 KV is covered. fp16 pools and K8V4 stay on the old kernel. - Tested on one gfx1201 card. gfx1200 compiles in but is untested. The gate is by architecture name. - Related open work: #313 adds RDNA3 (gfx11) WMMA kernels with a different fragment layout and touches the same file. A trial merge of this branch with #313, #312 and #270 completes without conflicts.
Sur le site
Liens install, modèles, releases.