Pull requests / #786

#786 hip: RDNA3 (gfx1100) WMMA kernel for the int8-KV QSA prompt attention (opt-in STRATA_HIP_WMMA, gfx12/gfx11)

closed · @xyzzing · 0 comments · View on GitHub

BenchmarksAMD / HIP

Description

## RDNA3 (gfx1100) WMMA kernel for the int8-KV QSA prompt attention — opt-in `STRATA_HIP_WMMA=1`, gfx12 + gfx11

This extends the S6 matrix-core prompt attention (currently gfx12-only) to RDNA3. gfx1100 keeps the ordered FP32 fallback as the default; the arm engages only under the existing `STRATA_HIP_WMMA=1` opt-in, and the default path is byte-identical to current behavior. (In particular, unlike #313, nothing in the default HIP path changes — the scalar fallback stays exactly as upstream ships it.)

Measured on an **RX 7900 XTX (gfx1100, 24 GB)**, engine 0.1.39, ROCm 10.2 nightly SDK (hipBLASLt 100500 table, gfx1100), production-mirrored config (spec 3, GR_V3, --spec-min-p 0.70, --kv int8, --kv-resident 32768). Since there is no AMD card in the upstream test ring, these numbers carry the PR.

### The layout that makes it work

RDNA3 runs the *same* `v_wmma_f32_16x16x16_f16` instruction as gfx12, but through the **unsuffixed** builtin — and its wave32 fragment layout differs from the gfx12-suffixed builtin's on **all three operands** (AMD [Matrix Instruction Calculator](https://github.com/ROCm/amd_matrix_instruction_calculator), `rdna3` tables, wave32, no modifiers; hardware-confirmed by a standalone probe before touching the kernel):

| operand | lane `l`, slot `i` holds |
|---|---|
| A | `A[l & 15][i]` |
| B | `B[i][l & 15]` |
| C/D | `D[(l >> 4) + 2*i][l & 15]` |

Packing the gfx12 way into the RDNA3 builtin passes compilation and fails parity at ~30× the output scale. The kernel's packing sites re-index accordingly; the online-softmax protocol, LDS staging, per-cell K/V scales and masking are unchanged (the B side still holds one cell per lane, so the per-cell scale folds exactly as on gfx12).

### Parity

Upstream's own `hip_prompt_attn_wmma` (first time it runs off gfx12) — 3/3 PASS:

| shape | vs FP64 (new) | vs FP64 (old FP32 kernel) | new vs old | speed |
|---|---|---|---|---|
| int8, ctx 32768, 2048 queries | 3.28e-06 | 2.13e-06 | 8.64e-06 | 34.05 → 8.41 ms/chunk (**4.05×**) |
| int8, ctx 1500, 1500 queries | 3.79e-06 | 1.78e-06 | 6.32e-06 | 12.32 → 2.41 ms/chunk (**5.12×**) |
| int8, ctx 2100, 256 queries | 4.45e-06 | 2.03e-06 | 5.44e-06 | 6.18 → 1.33 ms/chunk (**4.66×**) |

Objdump gate: exactly 32 `v_wmma_f32_16x16x16_f16` in the compiled kernel (the loop body — no fallback can masquerade).

### End-to-end (paired A/B, arms alternate, same binary, env-gated)

The FP32 fallback was measured at **25.6% (128K) / 26.9% (32K)** of prefill GPU time on this card — the single largest prefill phase.

| tier | arm-b / arm-a median | pair wins | prefill tok/s (a → b) |
|---|---|---|---|
| 128K | **+22.8%** | **5/5** (per-pair 1.212–1.253) | 1513.8 → 1863.2 |
| 32K | **+17.9%** | 5/5 | 1669.3 → 1968.0 |
| 1K decode | +2.3% (noise) | 3/5 | 65.9 → 67.4 tok/s |

Decode is flat at 32K (−0.03%) and 128K (medians equal) — the decode protection cap holds. Attribution agrees with the mechanism: the arm drops the qsa-attn share from 26–27% to 8.4–10.6% of prefill GPU time.

### Notes

- For #313's two review points (by @Niko1221): this PR value-parses the env (`e[0] == '1'`), and no `STRATA_WMMA_GFX11`-style compile define leaks into other device passes — the arch gate is the same runtime `gcnArchName` prefix check upstream already uses, extended to `gfx11`.
- The 2-file rebase surface (`qsa_prompt_attn.cu`, `qsa_prompt_attn_parity.cpp`) is deliberately minimal.
- Rebased on current main (6f32ec0, 0.1.39). Credit to @StevenChenSE's #313 for pioneering the gfx1100 WMMA prefill direction; the dense-GEMM half of that PR is orthogonal to this arm.

Related on strata.com

Editorial links to help you install, pick models, or read release notes — not part of the upstream thread.