onnxruntime
86af0f16 - [CUDA] Block-scaled MatMul decode GEMV: tensor cores, packed FP4 decode and M-tiling (#31155)

Commit
67 days ago
[CUDA] Block-scaled MatMul decode GEMV: tensor cores, packed FP4 decode and M-tiling (#31155) ### Description Five optimizations to the decode GEMV fast path of the `MatMulBlockQuantizedFp4Weight` and `MatMulBlockQuantizedFp8Weight` contrib ops. All are behind existing kill switches (`ORT_FP4_GEMV_MMA=0`, `ORT_FP8_GEMV_MMA=0`, `ORT_FP4_GEMV_ROW_TILING=0`) for A/B. | commit | change | |---|---| | `fp8 blocked scaled gemm` | hoist the fp32 widening out of the inner loop when `RowsPerWarp > 1` | | `prmt quad decode` | decode four E2M1 codes per `prmt.b32` instead of one | | `FP4 dense GEMV: grid-gated row tiling` | tile over M, gated on the column-block count | | `Run FP8 weight-only decode GEMV on tensor cores` | `mma.m16n8k16` | | `Run the NVFP4 decode GEMV on tensor cores` | `mma.m16n8k16` | ### Motivation and Context The shipped GEMV kernels were tuned for `M = 1` (plain autoregressive decode). Speculative / MTP decoding makes the verify forward `M > 1` wide, and at `M = 4` the `M = 1` tuning inverts: * cuBLAS fp16 gets *faster* per byte (3.1 TB/s — it amortizes the weight read over 4 rows); * the FP8 GEMV gets *slower* per byte (2.35 -> 1.25 TB/s), because it re-widens B to fp32 once per row and A once per column, so it goes ALU-bound. Measured on H200 (µs, shipping `<R,1,1>` kernel vs cuBLAS) — flipping dense projections to FP8 without touching the kernel was a *regression* at `M = 4`: | shape (NxK) | M=1 cuBLAS / fp8 | M=4 cuBLAS / fp8 | |---|---|---| | in_proj_qkv 8192x2048 | 10.9 / **7.1** | **10.9** / 13.7 | | in_proj_z 4096x2048 | 8.3 / **4.8** | 8.3 / **8.1** | | out_proj 2048x4096 | 8.8 / **5.1** | **8.3** / 9.1 | All numbers below: 1x H200 SXM (SM90, 132 SM, ~4.8 TB/s HBM), CUDA 13.0, Qwen3.6-35B-A3B-NVFP4 + MTP `N=3` (verify batch `M=4`). ### 1. FP8: hoist the fp32 widening for `RowsPerWarp > 1` — −0.13 ms/step `M == 1` keeps the original scalar path (hoisting measured slightly *worse* there). For `R > 1`, B is widened once per `(col, half)` and reused by all rows, A once per `(row, half)` and reused by all cols; the 16-element chunk is consumed in two 8-element halves to cap live registers. Instructions per 16 weight bytes: `8 + 48*R` -> `8*C + 16*R + 16*C + 16*R*C` (200 -> 120 per column at `R=4, C=2`). **The fma order is unchanged, so the output is bit-identical.** After the fix, `M=4` FP8 beats cuBLAS by 1.04–1.59x and the old `<R,1,1>` by 1.13–1.41x. End-to-end 2x2 (graph ON, 400 steps, 64 warmup, 2 reps), ms/step: | | old kernel | new kernel | |---|---:|---:| | v10 (fp16 dense) | **10.80** | 10.711 | | v11 (FP8 dense) | 10.90 | **10.674** | **10.80 -> 10.674 = −0.13 ms/step (−1.2%).** Acceptance rate unaffected (2.5–2.75 tok/step in every cell). ### 2. FP4: `prmt` quad decode — −0.26 ms/step (−2.4%) `prmt.b32` is a 4-byte table lookup in one instruction, so a whole 32-bit word (8 E2M1 codes) is decoded at once instead of per element. | shape | M | scalar µs | prmt µs | speedup | |---|---|---:|---:|---:| | lm_head 248320x2048 | 1 | 186.97 | **157.03** | 1.19x | | lm_head | 2 | 368.90 | **310.01** | 1.19x | | lm_head | 4 | 733.17 | **613.70** | 1.19x | | shared gate/up 512x2048 | 4 | 4.76 | **4.22** | 1.13x | | shared down 2048x512 | 4 | 6.01 | **5.14** | 1.17x | SASS instruction count: half 352 -> 288 (−18%), bf16 376 -> 336 (−11%). End-to-end (graph ON, 400 steps, 3 interleaved reps): 10.766 -> **10.504 ms/step**; every prmt rep beats every baseline rep. **Bit-identical** — an exhaustive 256-byte sweep against the old `Raw()*Scale()` path gives 0 mismatches for both half and bf16. ### 3. FP4: grid-gated row tiling — −0.14 ms/step (−1.4%) Process `RowsPerBlock` rows of Y at once so the `uint4` weight load, the `prmt` decode and the scale load are amortized across rows. Gated on the column-block count, because at small N the grid is already narrow and tiling *loses*: | shape | N | col blocks | untiled µs | tiled µs | speedup | |---|---|---|---:|---:|---:| | lm_head | 248320 | 31040 | 615.8 | **537.7** | 1.15x | | shared down | 2048 | 256 | 4.30 | **3.33** | 1.29x | | shared gate/up | 512 | 64 | 3.47 | 4.27 | 0.81x (gate keeps this untiled) | FP4 GEMV family (graph OFF): 1.099 -> 1.066 (launch bounds) -> **0.950 ms/step** (1.157x). End-to-end (graph ON, 200 steps, 4 interleaved reps): 10.448 -> **10.306 ms/step**. **Bit-identical** (per-row fp32 accumulation order unchanged). ### 4. FP8 on tensor cores (`mma.m16n8k16`) — −0.27 ms/step (−2.7%) | shape | M | FMA µs | MMA µs | speedup | |---|---|---:|---:|---:| | 8192x2048 | 1 | 6.3 | **5.1** | 1.24x | | 8192x2048 | 4 | 9.8 | **5.2** | 1.88x | | 8192x2048 | 8 | 17.5 | **5.7** | 3.07x | | 4096x2048 | 1 | 4.8 | **4.0** | 1.20x | | 4096x2048 | 4 | 6.9 | **4.1** | 1.68x | | 2048x4096 | 1 | 5.1 | **4.2** | 1.21x | | 2048x4096 | 4 | 8.0 | **4.4** | 1.82x | FP8 GEMV family (graph OFF, 130 launches/step): **1.052 -> 0.713 ms/step (1.48x)**; all kernels 7.368 -> 7.021 ms/step. End-to-end **9.80 -> 9.54 ms/step**. Now at 1.28 GB/step = 1.8 TB/s = 37% of HBM peak. Not bit-identical, but not worse: E4M3->FP16 is lossless and the `f16 x f16` products are exact in fp32, so max error against an FP64 reference is **identical for both kernels (2–4e-4)**. ### 5. FP4 on tensor cores (`mma.m16n8k16`) — −0.55 ms/step (−5.8%) Largest single win of the campaign. | shape | launches/step | scalar µs | MMA µs | speedup | |---|---|---:|---:|---:| | lm_head N=248320 K=2048 | 1 | 0.5376 | **0.1083** | **4.96x** | | gate/up N=512 K=2048 | 80 | 0.2786 | 0.2391 | 1.17x | | down N=2048 K=512 | 40 | 0.1327 | 0.1007 | 1.32x | | **FP4 GEMV family total** | 121 | **0.9488 ms** | **0.4480 ms** | **2.12x** | All kernels: 7.018 -> 6.526 ms/step. End-to-end (graph ON, 5 interleaved reps): 9.768 / 9.814 / 9.824 -> 9.558 / 9.584 / 9.466, i.e. **9.54 -> 8.99 ms/step**. **Bit-identical** (E2M1 decode order and fp32 accumulation order unchanged). The kernel now sits at 56% of HBM peak, and two further optimizations measured 0%, so this path is considered closed. ### Tests * `onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc` * `onnxruntime/test/contrib_ops/matmul_block_scaled_fp8_test.cc` — including `GemvSpeculativeDecodeTilesFp16` (`m` in {2,3,4} x `n` in {1026,2050,4098,8194}, ragged tails, row- and column-varying data). ### Methodology note All end-to-end deltas are quoted as **ms/step** from a fixed-step measurement, never tok/s: any numerics change alters the generated sequence and therefore the MTP acceptance rate, which swamps the speed delta. Per-kernel durations are taken with CUDA graphs **off** — `nsys --cuda-graph-trace=node` inflates durations ~35% globally and up to 3.8x for large-grid kernels (FP4 lm_head: 2840 µs reported vs 741 µs real).
Author
Parents
Loading