Pull requests / #313

#313 RDNA3 WMMA kernels for the gfx1100 prefill (opt-in, runtime-gated on gfx11; rebased on 0.1.30)

closed · @StevenChenSE · 0 comments · View on GitHub

BenchmarksSetup & installAMD / HIPNVIDIA / CUDADocumentation

Description

Second attempt of #260, rebased on current **main (engine 0.1.30)** — reworked per the review. Single commit, 10 files.

## Addressing the review

| review point | what changed |
|---|---|
| opt-in (`STRATA_WMMA_GEMM=1`) | done — nothing runs unless asked. `STRATA_WMMA_GEMM=1` enables the dense WMMA GEMMs (fp16 + bf16), `STRATA_PA_WMMA=1` the WMMA prompt attention; `STRATA_WMMA_BF16=0` still excludes just bf16. With the switches **unset, WMMA does not run at all** — the default hipBLASLt-table-to-fallback path is unchanged and output is bit-identical to today's build. When opted in, WMMA dispatches ahead of the table; measurement below shows that arrangement is also the faster one on this card. |
| runtime gate on gfx11 only | done — both paths query `hipGetDeviceProperties().gcnArchName` and return false without a `gfx11` prefix, so a gfx1201, a mixed-arch build, or any other device never executes WMMA code regardless of what was compiled. Additionally the `STRATA_WMMA_GFX11` compile definition is now set only for gfx11 `CMAKE_HIP_ARCHITECTURES` configurations. |
| fp16 rounds toward zero | acknowledged and bounded: the matrix core accumulates in fp32 with RZ while hipBLAS rounds to nearest. With the switches off this is unobservable by construction; with them on, the parity test bounds the deviation at 4 ULP (every measured delta ≤ 1 ULP and toward zero; bf16 compares exactly), documented in `docs/AMD_HIP.md` and `tests/hip/prefill_wmma_gemm_parity.cpp`. |
| CMakeLists.txt BOM / comment encoding | preserved — patched at byte level; BOM re-verified present after the edit (`efbbbf`), and `cmake/hip_backend.cmake` (no BOM upstream) likewise untouched in encoding. |
| with/without numbers | below — including against a **freshly calibrated hipBLASLt table**, not just the version-mismatched fallback. |

## What the PR contains

* `src/prefill/wmma_gemm.{cu,h}` — dense GEMM over the RDNA3 WMMA fragment layout (16x16x16, wave32 "doubled" inputs), fp16 + bf16, dispatched from `Gemm::f16`/`Gemm::bf16`.
* `src/kernels/cuda/qsa_prompt_attn.cu` — the prompt attention on WMMA instead of the ordered FP32 fallback.
* `tests/hip/prefill_wmma_gemm_parity.cpp` — ctest-registered parity vs a double-precision host reference, 440 shapes × 4 beta/ldy configurations × 2 dtypes (partial tiles, padding, guards, declined shapes); passes on this base (freshly rebuilt from this branch's sources).
* Minimal build wiring: one source in `strata_prefill`, the arch-gated compile definition, the ctest registration, and the `CUDART_INF_F` mapping hip_compat needs.
* A "RDNA3 WMMA kernels" section in `docs/AMD_HIP.md`.

## With/without numbers

One RX 7900 XTX, 6-core host, PCIe 4.0 x16, engine 0.1.30, the documented measured configuration (`--prefill 8192 --spec 4 --kv int8 --kv-resident 32768 --adapt-every 0 --pcie-frac 0`, `STRATA_PREFILL_MMQ=1`), two interleaved repetitions per cell.

**Against the strongest BLAS available here.** The shipped tables (100100/100200) don't match this ROCm's hipBLASLt 1.4.1 (100401), so I calibrated a table for this runtime with the repository's own `tools/hip/tune_hipblaslt` and compared against it (4K, ranking consistent in both pairs):

| arm | prefill (tok/s) | decode (tok/s) |
|---|---:|---:|
| WMMA (`GEMM=1 PA=1`) | **1083.7 / 1308.4** | 32.1 / 41.3 |
| calibrated hipBLASLt table (WMMA off) | 852.2 / 1176.6 | 38.4 / 40.7 |
| neither (hipblasGemmEx fallback) | 658.6 / 552.8 | 37.3 / 37.3 |

WMMA leads the **calibrated** table by 10–27 % on prefill in both pairs, and the plain fallback by 55–90 %; decode is a wash at this sample size.

**Switch sweep** (the "off" arm here is the version-mismatched hipblasGemmEx fallback):

| tier | switch state | prefill (tok/s) | decode (tok/s) |
|---|---|---:|---:|
| 1K | off | 376 / 455 | 35.2 / 45.4 |
| 1K | `WMMA_GEMM=1` | 471 / 551 | 39.7 / 51.5 |
| 1K | `GEMM=1 PA=1` | 619 / 471 | 50.1 / 50.3 |
| 32K | off | 790 / 800 | 57.3 / 55.8 |
| 32K | `GEMM=1 PA=1` | **1501 / 1484** | 55.1 / 57.1 |

Prefill: **+31 %** at the 1K median, **+88 %** at 32K vs the fallback; decode unchanged within this host's variance. Host and configuration named because they matter: expert offload over PCIe and a 6-core CPU are part of the setup.

## Integration notes (without these the kernels compile but never run)

1. this tree's HIP macro is `STRATA_USE_HIP` — a private `STRATA_BACKEND_HIP` silently compiled the dispatch out;
2. the host pass of a HIP compile does not define `__gfx1100__` (it defines `__HIP_DEVICE_COMPILE__`), so the build supplies `STRATA_WMMA_GFX11`;
3. `hip_compat` needs `CUDART_INF_F` (CUDA's `math_constants.h` spelling) mapped to `HIP_INF_F`.

## Not claimed

* Decode gains: decode is unchanged within this host's variance in every comparison above.
* One host, one model pack, two repetitions per cell. With the switches off, behavior and output are unchanged from main.

Related on strata.com

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