Pull requests / #1006
#1006 HIP RDNA2 (gfx103x): run the prompt path's 16-bit GEMMs as FP32 SGEMMs (RX 6800: prompts 312 -> 452 tok/s)
closed · @vallicgrr · 0 commentaires · Sur GitHub
BenchmarksSetup & installServer & APIAMD / HIPNVIDIA / CUDAModels & quantsDocumentationWindowsLinux
Description
On gfx1030 the rocBLAS that setup installs (ROCm 10.2.0a20260930) has tuned kernels for FP16 -> FP16, int8 and FP32 -> FP32 only. The prompt path's dense GEMMs take FP16 or BF16 inputs and write FP32 (rocBLAS's HS / BS types), so every one of them runs a fallback kernel at about 5 TFLOPS. This PR widens both inputs to FP32 on gfx103x and runs SGEMM instead. - `Gemm::rdna2_sgemm` (`src/prefill/gemm.cu`): the weight is widened once per call, the activations in slices of up to 128 MiB of FP32. Widening FP16 / BF16 to FP32 is exact, and SGEMM accumulates in FP32 like the native call. - `Gemm::bf16` and `Gemm::f16` try it after hipBLASLt and before the native `hipblasGemmEx`. `Gemm::native` goes through `f16`, so dequantized weights take the route too. - It is on by default for gfx103x only. Every other AMD architecture and CUDA are unchanged. N < 64 (routers, gates) keeps the native call, because there the widening costs more than the GEMM gains. If the FP32 buffers cannot be allocated, the native call runs. - `STRATA_RDNA2_SGEMM=0` / `=1` turns it off / on for any AMD card. These are the A/B arms below. - `docs/AMD_HIP.md`, RDNA2 section: the measurement, and a Windows gfx1030 report. **Measured** on an RX 6800 16 GB over OCuLink (PCIe 4.0 x4, 7.1 GB/s), Ryzen 7 8845HS, 28.8 GB RAM, Windows 11. Coder IQ1_M with setup's arguments (64K context, MTP), fresh 12K-token prompts with no reuse, 4 per arm, the same binary with `STRATA_RDNA2_SGEMM=0` against the default: | | native GEMMs | SGEMM route | |---|---|---| | prompt | 312 tok/s | 452 tok/s | | prompt GPU timeline (`STRATA_PREFILL_TIMING`) | 37.2 s | 25.5 s | | gdn / qsa proj / hc read | 9.4 / 3.6 / 3.85 s | 2.8 / 1.4 / 2.84 s | | decode | 29-34 tok/s | 29-35 tok/s | - **Per call** (the engine's exact `hipblasGemmEx`: opA = T, opB = N, FP32 out, fp32 compute): - N 10240 x T 8192 x K 2560: 86.7 -> 27.7 ms, including the widening. - N 2560 x T 8192 x K 6144: 53.2 -> 17.9 ms. - N 320 x T 8192 x K 10240 (BF16): 10.9 -> 7.5 ms. - Summed over one prompt's FP16 calls: 14.1 -> 4.7 s. - **Accuracy:** I checked the native call and the SGEMM route against a float64 reference, using random inputs in [-1, 1] and 5 shapes in each of FP16 and BF16 at T = 2048. - The two routes differ by at most 4.3e-6 of the largest output. - The SGEMM route is as close to the reference as the native call, or closer, on all 10 shapes. - **End to end:** a 9-task code-repair benchmark, where the model patches seeded bugs in a browser game until 38 browser checks pass. Both arms solved 9/9 on the first try; model time went from 458 s to 374 s. - **Builds:** `tools/hip/build_windows.bat` with `STRATA_HIP_ARCHS=gfx1030`, using the `gfx103X-all` 10.2.0a20260930 wheels. The engine starts through `serve/server.py` and serves both APIs. - **Test:** `tests/hip/prefill_gemm.cpp` gains two cases that cross the route's activation slice boundary (T 3300, K 10240): FP16 with `beta = 1` and a padded `ldy`, and BF16. All six cases pass against the float64 reference on the RX 6800, with rel_l2 at most 1.8e-6. - **ctest** (`build_windows.bat tests`, the RX 6800 with the 780M hidden by `HIP_VISIBLE_DEVICES=1`): 59 passed, 2 skipped (gfx12 only), 6 failed out of 67. A build of `main` (6f32ec0) fails the same six tests with the same messages, and so does this branch with `STRATA_RDNA2_SGEMM=0`: - `ple_parity`, `expert_parity` and `pool_test` need model fixtures (`Q2_0/...gguf`, `pack/full/experts.bin`) that are not on this machine. - `expert_cache_segmented_test`: `--vram-elastic` is CUDA-only. - `hip_handoff`: "separate copy/ring timeout". - `hip_prefill_mmq_parity`: "synthetic-Q2_0-GU-pass0: non-finite or unwritten MMQ output". - The last two look like gfx1030 (or Windows) issues of their own. I have not looked into them. - **Not tested:** Linux, and other gfx103x cards. **Also noticed, not changed here:** after this change, the prompt path's biggest phase on this PC is waiting for streamed experts (`wait copy`, ~6 s of 25.5 s), and then the QSA attention fallback (3.7 s).
Sur le site
Liens install, modèles, releases.