Pull requests / #1007

#1007 hip: the QSA block scores on rocBLAS SGEMM, the solution measured on the card (opt-in; a 128K prompt +2.8% on gfx1030)

closed · @xjc10 · 0 comments · View on GitHub

BenchmarksMulti-GPUAMD / HIPNVIDIA / CUDAModels & quants

Description

The QSA block scores - for every (query, KV block) the four indexer heads' relu(q . k) summed, K = 128 - run on sm_80+ as one GEMM (3xTF32 MMA) and on gfx12 as a WMMA GEMM (opt-in); every other HIP card runs the warp kernel, one warp per (query, block). On an RX 6900 XT (gfx1030, no matrix cores) that kernel is the prompt path's one term that grows with the context: 27 ms per 8K-token chunk at the start of a prompt, 752 ms per chunk at a 128K context (12 QSA layers x 32 batches of 256 queries x ~2 ms), 6.3 s of a 128K prompt's 66 s.

This PR gives those cards the same GEMM through rocBLAS SGEMM, opt-in (`STRATA_SELECT_SGEMM=1`):

- Per batch of queries and indexer head, `rocblas_gemm_ex` computes `C[reach x nq] = pooled[reach x 128] . q_h^T` into a scratch slab (nq x max_blocks floats, allocated once per device), a kernel adds relu(C) into the scores, and the tail block n_bid is scored by `block_scores_tail_kernel` as the tensor-core paths do. FP32 throughout, rocBLAS's summation order: FP32-level and not bitwise the warp kernel, the class of accuracy the sm_80 path has.
- rocBLAS's own kernel choice for this skinny shape is slow on gfx1030 (reach = 32768, nq = 256: 1.56 ms, 1.4 TFLOPS, while the best of its 244 listed solutions runs it in 0.11 ms, 14x), and the solution indices are a property of the rocBLAS build. So the scorer measures: on the first call it times, for each reach bucket (powers of two up to the capacity), every solution `rocblas_gemm_ex_get_solutions` lists once and the eight fastest three times, and keeps the fastest listed one - never rocBLAS's own choice, even when that is close at the bucket's top: a solution's time falls with the reach, the heuristic's choice changes with the exact shape (reach 65,536: 0.22 ms; reach 32,769, the same bucket: 1.6 ms). rocBLAS's choice is the fallback for a shape the solution refuses. The measurement is at the batch size of that call (the prompt path's full batch); a chunk's last, smaller batch reuses it (measured within 1.3x of that size's own best and a quarter of rocBLAS's choice), a larger batch measures again. It costs about 0.5 s per card, once per engine start, inside the first prompt's first chunk. `STRATA_SELECT_SGEMM_TUNE=0` keeps rocBLAS's choice; `STRATA_SELECT_SGEMM_VERBOSE=1` prints each bucket's measurement.
- CMake: `find_package(rocblas)` (quiet) when hipBLAS is 1.0 or newer sets `STRATA_ROCBLAS_AVAILABLE`; `strata_kernels` links `roc::rocblas` then. Builds without the package keep the warp kernel.

**Measured** (RX 6900 XT 16 GB, gfx1030, ROCm 10.0.0, rocBLAS 5.6.0; `6f32ec0` + this):

`qsa_select_bench` (synthetic, the repository's scorer-vs-scorer comparison and FP64 gate), 256 queries:

| context | warp kernel scores | SGEMM scores | FP64 gate | selections |
| --- | ---: | ---: | --- | --- |
| 131,072 (32,769 blocks) | 2.739 ms | **0.811 ms** (3.4x) | PASS, max err 5.6e-05 at scale 268 | identical 256/256, cells differing 0.0000% |
| 32,768 (8,193 blocks) | 0.679 ms | **0.304 ms** (2.2x) | PASS, max err 5.5e-05 at scale 269 | identical 256/256 |

End to end (2x RX 6900 XT layer split, Qwen3.8-Flash-Next GSQ-RCO IQ3_S, `--kv int8 --kv-resident 32768 --max-context 131072 --adapt-every 100000`, temperature 0, 256 generated tokens, one run per arm, two engine starts; `STRATA_PREFILL_TIMING=1`; `main` has no rocBLAS table (#981), so these are `main`'s prompt speeds, not this rig's best):

| prompt | `STRATA_SELECT_SGEMM` unset | `=1` | `qsa select` per 8K chunk at the prompt's end | the prompt's whole `qsa select` |
| --- | ---: | ---: | --- | --- |
| 50,517 tokens | 756 tok/s (66.8 s) | 753 tok/s (67.1 s) | 282 -> 115 ms | 1,006 -> 1,123 ms (1,090 of it the measurement) |
| 130,681 tokens | 810 tok/s (161.3 s) | 833 tok/s (156.9 s) | 764 -> 271 ms | 6,301 -> 2,445 ms |

The selection scores are the one prompt term that grows with the context, so this is a long-context change: per 8K chunk the GPU timeline grows from 8.29 s (first chunk) to 9.42 s (the 16th) without it and from 8.42 to 8.95 s with it; at 50K the gain is about the size of the one-time measurement, below 16K there is nothing to gain. The larger part of what is left in `qsa select` at 128K is top-k, which this PR does not touch.

Decode is unchanged (the decode path scores one query on the warp kernel and is not touched): 55-56 tok/s in both arms. The generated text is the same in both arms; the drafts accepted differ by a few (near-tie selections), as between the warp kernel and the sm_80 scorer.

Limits: one machine, one card family, one run per arm. The measurement picks by time alone; its result could differ between engine starts when two solutions tie, which changes near-tie selections, not the output's class of accuracy. Independent of #835 / #981 (the same rocBLAS kernel-choice problem on gfx1030, there with an offline table for the prompt path's fixed shapes; here the shape varies per call, so it is measured at run time - happy to move to one mechanism if preferred); the two add the same `find_package(rocblas)` lines to `cmake/hip_backend.cmake` with different comments, so whichever lands second needs a one-line rebase there. The relu-and-add pass runs once per head (0.7 of the 0.8 ms at 128K is these four passes and the tail kernel, the SGEMMs are 0.11); folding the four heads into one pass needs four scratch slabs (134 MB at a 128K context) and is a follow-up.

Developed with an AI coding assistant; every number above was measured on 2x RX 6900 XT (gfx1030, PCIe 4.0 x8 each) / Ryzen 5 5600X, ROCm 10.0.

Related on strata.com

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