Pull requests / #337
#337 hip: gfx12 (RDNA4) QSA select - WMMA block scorer, prompt-aware top-k dispatch
closed · @bsorensen110 · 0 コメント · GitHub で見る
BenchmarksMulti-GPUAMD / HIPNVIDIA / CUDAModels & quants
本文
## What Two changes to the **prompt path's QSA selection** (`src/kernels/cuda/qsa_select.cu`) for AMD RDNA4 (gfx1200/gfx1201), measured on a Radeon AI PRO R9700 (gfx1201, 32 GB) serving Swift 1.5 IQ3_XXS with `--kv int8 --max-context 262144`. 1. **`qsa_block_topk` dispatch follows the prompt, not the capacity.** The register kernel was chosen by `max_blocks`, which is a *capacity* (`--max-context / 4 + 2`). With a 262144-cell capacity every prompt, even a 4K one, took the slow reference top-k (up to ~10x slower at 128K+). The prompt path now passes the batch's active block count (the argument `qsa_block_scores` already takes). Decode passes none and keeps the capacity rule and the original register width, so **decode is unchanged**. On HIP the register kernel is also instantiated for 66 blocks per thread (up to 270,336 cells, covering a 262144 `--max-context`), and below ~28K cells, where it measured 0.2-0.9x of the reference on gfx1201, the reference is kept. This applies to CUDA too (the new argument is optional and the width is unchanged there), but I could only measure on gfx1201. 2. **`qsa_block_scores_tc` on gfx1200/gfx1201: a WMMA block scorer** (`v_wmma_f32_16x16x16_bf16`, wave32). Each operand is split into three bf16 slices (hi / mid / lo are exact 8-bit mantissa slices of the fp32, no scaling needed) and the six products of order <= 2 are summed. The hi*hi chain keeps its own accumulator and the five corrections another, added once at the end: a single shared accumulator rounded the corrections at the big accumulator's ulp (8e-7 of the score scale, failing the harness gate at short context); split, it is 1.6e-7, the warp kernel's own error is of the same order. The tail block is the warp kernel's arithmetic. Like the CUDA TF32 kernel it is **not bitwise** with the warp kernel (a near-tie can select differently). Gated at compile time (`__gfx1200__`/`__gfx1201__`) and run time (`gcnArchName`); other HIP targets keep the warp kernel; CUDA is unchanged. `qsa_select_bench` now builds on HIP, takes the engine's capacity as a 4th argument, checks the scorer against an FP64 reference (gate: no worse than 4x the warp kernel's error, floored at 1e-6 of the score scale), exits non-zero on a failed gate, and is registered with ctest (32K context under a 262144-cell capacity). Independent of #329 (different kernel, disjoint files apart from the CMake block); a merge dry run against `main`, #322 and #329 is clean. ## Results (R9700, same binary, only `STRATA_SELECT_OLD=1 STRATA_TOPK_OLD=1` differs between arms) Engine's own prefill timing, service stopped, prompts that were not partly cached are compared by throughput: | prompt | old | new | speedup | `qsa select` phase | |---|---|---|---|---| | 107.7K tokens | 79.5 s | 59.5 s | 1.34x | 23.3 s -> 3.4 s | | 188K tokens | 1,220 tok/s | 1,812 tok/s | 1.48x | 57.7 s -> 7.7 s | | 245K tokens (16K reused) | 1,020 tok/s | 1,748 tok/s | 1.71x | 115.6 s -> 15.5 s | | 262K tokens | 1,019 tok/s | 1,747 tok/s | 1.71x | (see note) | | 30.9K tokens | 17.0 s | 16.6 s | 1.02x | 0.68 s -> 0.26 s | | <= 8.8K tokens | | | ~1.01x | small share | Note on the 262K row: the old arm's second 262K request reused a 16K prefix and the new arm's did not, so wall seconds are not comparable there (149.5 s vs 240.5 s); I compared tokens/s. Every other prefill phase is within ~1% between arms, and an old-vs-old control run matched the old arm within 0.1% on wall time and each phase. Standalone bench, 256 queries, capacity 262144: scores 0.81 -> 0.18 ms at 32K, 3.5 -> 0.65 ms at 131K, 6.8 -> 1.5 ms at 262K; top-k 3.9 -> 0.51 ms at 131K, 12.9 -> 1.3 ms at 262K. Top-k ids equal the reference's in 256/256 queries at every size tested. Scores vs the warp kernel at 262K: 254 of 256 queries select identically, 0.0011% of cells differ (near-ties). ## Correctness evidence, and what is missing - Needle-in-a-haystack: 32K (depths 10/50/90), 124K, 188K (depth 90), and 262K (depths 10 and 50, 261.7K tokens each) found in the new arm; the same needles are found in the old arm. - ctest 43/43 on this branch (47 registered, the usual four excluded: `ple_parity`, `platform_memory_test`, `expert_parity`, `pool_test`), including `qsa_select_bench`. - **No teacher-forced top-1/logit comparison against the old kernel was run.** Reply *text* differs between arms, but it also differs between two identical old-config runs (3 of 9 prefill-probe replies identical), so text comparison cannot discriminate here. - Measured on gfx1201 only. The CUDA path is unchanged apart from the optional `active_blocks` argument, which is not measured on NVIDIA. Tested with ROCm 10.2.0a nightly (gfx120x), `-DCMAKE_HIP_ARCHITECTURES=gfx1201`.
関連リンク
インストール・モデル・リリースへの站内リンク。