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 Kommentare · Auf GitHub
BenchmarksSetup & installAMD / HIPNVIDIA / CUDADocumentation
Beschreibung
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.
Mehr auf der Site
Links zu Install, Modellen, Releases.