Pull requests / #368
#368 QSA select: RDNA3 (gfx1100) WMMA block-scores arm on #337's arch dispatch (opt-in STRATA_SELECT_WMMA=1)
closed · @xyzzing · 0 commentaires · Sur GitHub
BenchmarksAMD / HIPNVIDIA / CUDA
Description
# QSA select: RDNA3 (gfx1100) WMMA block-scores arm on #337's arch dispatch (opt-in STRATA_SELECT_WMMA=1)
Rebase of #312 as requested ("rebase this onto #337's HIP scorer entry point (it dispatches by arch) and keep it
opt-in under STRATA_SELECT_WMMA=1"). Stacked on #337: it merges that branch first, then adds the gfx1100 arm.
## What the dispatch looks like now
`qsa_block_scores_tc` on HIP dispatches by architecture, one arm per card family:
1. **gfx12 (RDNA4)** — #337's `block_scores_wmma_kernel` (bf16 three-way split), unchanged.
2. **gfx1100 (RDNA3)** — this PR's `block_scores_wmma_gfx11_kernel`: rocWMMA 16x16x16 FP16 operands, FP32
accumulation, wave32; keys and queries converted to FP16 per call into an internal scratch (the conversion is
part of the arm's cost). **Opt-in: `STRATA_SELECT_WMMA=1`** — its scores are FP16-input with another summation
order, so they can change selection. `STRATA_SELECT_OLD=1` still forces the warp kernel above all of this.
3. anything else — `return false`, the warp kernel (unchanged default behavior; no card sees a new arm by default).
The arm that ran is printed once (`strata select: prompt scorer arm = ...`), so an opt-in that quietly fell back
cannot masquerade as the fast path.
## Evidence (all measured on the card the PR targets)
- **Parity, 5/5 cases incl. adversarial fixtures** (re-run on this branch's own 0.1.31 build) (`qsa_select_wmma_parity`, registered with ctest): the warp arm
is bitwise == its host model; the tail block is bitwise == the warp arm (shared tail kernel — no new tail
arithmetic); every score inside the pre-derived bound `B(qi,b) = 2^-9.5 * S(qi,b) + 8 * 2^-24 * |score|` vs a
FP64 oracle (worst utilisation 0-11%); selection ids agree wherever the gap > 2B; the +2.5B negative control
correctly FAILS the suite (the harness can detect a defect of the size the bound claims to cover).
- **Kernel bench @128K ctx (nq=256)**: warp 2.735 ms → wmma 0.235 ms = **11.64x** on the kernel, re-measured on this branch's 0.1.31 build (the 0.1.24-base campaign measured 11.05x — same kernel, new base); the
same kernel measured paired end-to-end at **+8.7% at 128K** (1K +0.2% ... 64K +3.8%; 123 paired cells,
`STRATA_SELECT_WMMA=1` vs default; decode within the protected cap) - that full table is in #312's description
and carries over: the kernel is byte-identical through the 0.1.29/0.1.30/0.1.31 rebases, only the base moved.
- The bench (`qsa_select_bench`) gains an FP64 accuracy gate sampled against all present arms and returns non-zero
on a failed gate or a register-top-k mismatch (addressing the Copilot review items on #337's bench for our arm).
## Files
`CMakeLists.txt` (parity + bench registration, rocWMMA detection), `include/strata/kernels/qsa_select.hpp`,
`src/kernels/cuda/qsa_select.cu`, `src/kernels/qsa_select_bench.cpp` (merged with #337's capacity-arg + FP64-gate
form), `src/kernels/qsa_select_wmma_parity.cpp`, `src/prefill/prefill.cpp` (call site unchanged from #337's —
the dispatch is internal to `qsa_block_scores_tc`).
We have no RDNA4 card, so the gfx12 arm is #337's and theirs to carry; on our RX 7900 XTX every number above is
personally reproduced. Default behavior on every architecture is unchanged unless `STRATA_SELECT_WMMA=1`.
---
Validated on this branch's own 0.1.31 build (gfx1100): parity 5/5, bench 128K 11.64x kernel with the wmma arm at worst 0.09 of its pre-declared bound, register top-k ids identical 256/256. The bench's FP64 gate applies each arm's own declared contract (4x-warp for the FP32-level arms; the 2^-9.5*S fp16-conversion bound for this arm - the same bound the parity suite enforces).
Sur le site
Liens install, modèles, releases.