[CUDA] MatMul: add an opt-in split-K GEMV for small-N fp16 shapes - #31478
Conversation
There was a problem hiding this comment.
Pull request overview
Adds an experimental, opt-in CUDA fast path for decode-time MatMul shapes that behave like small‑N GEMVs (fp16, row-major, no transpose), using a split‑K kernel to increase SM utilization. The path is gated behind ORT_ENABLE_SMALL_N_GEMV and falls back to the existing cuBLAS implementation by default.
Changes:
- Introduces
SmallNGemvSplitKKerneland host-side helpers (eligibility, split‑K heuristic, workspace/counter sizing, launch). - Adds a guarded fp16 fast path in
MatMul<T>::ComputeDefaultbefore the cuBLAS GEMM call. - Adds a CUDA unit test that directly launches the kernel and checks counter reset behavior across repeated launches.
Reviewed changes
Copilot reviewed 4 out of 4 changed files in this pull request and generated 1 comment.
| File | Description |
|---|---|
| onnxruntime/core/providers/cuda/math/matmul.cc | Adds env-var gate and routes eligible fp16 MatMul into the new small‑N GEMV split‑K kernel. |
| onnxruntime/core/providers/cuda/math/matmul_small_n_gemv.h | Declares eligibility + sizing helpers and the launch entry point for the split‑K GEMV kernel. |
| onnxruntime/core/providers/cuda/math/matmul_small_n_gemv.cu | Implements the split‑K kernel and dispatch for M=1..8, including deterministic slice-ordered reduction. |
| onnxruntime/test/providers/cuda/test_cases/matmul_small_n_gemv_test.cc | Adds direct-kernel unit tests covering M variants, tile boundaries, and counter reuse/reset semantics. |
Suppressed comments (1)
onnxruntime/core/providers/cuda/math/matmul.cc:348
- The kernel resets the arrival counter internally (matmul_small_n_gemv.cu sets counter[blockIdx.x] = 0), and the unit test relies on that by launching twice without re-memsetting. Here the MatMul fast-path still does a cudaMemsetAsync on every launch, which adds overhead on the microsecond-scale GEMV path and defeats the intended "clear once, then reuse" design described in the PR.
Consider caching a counter buffer across launches (or otherwise ensuring one-time initialization) so the per-call memset can be removed, or alternatively drop the in-kernel reset if you intend the counter to be per-launch scratch only.
const size_t counter_elements = SmallNGemvCounterElements(n);
auto counter = GetScratchBuffer<unsigned int>(counter_elements, this->GetComputeStream(ctx));
CUDA_RETURN_IF_ERROR(cudaMemsetAsync(counter.get(), 0, counter_elements * sizeof(unsigned int), Stream(ctx)));
auto workspace = GetScratchBuffer<float>(SmallNGemvWorkspaceElements(m, n, k), this->GetComputeStream(ctx));
|
Thanks for putting this behind an opt-in flag and documenting the real-model benchmark result. The split-K decomposition, bounds, fixed-order reduction, build integration, scratch allocation, and stream ordering all look reasonable. However, I think this needs changes before merge. Blocking issue: cross-block workspace visibilityIn CUDA's documented last-block reduction pattern states that a memory fence orders operations but does not, by itself, guarantee visibility to other blocks. The communicated buffer must use volatile/cache-bypassing accesses so the final block cannot consume stale L1 data. Here This is particularly relevant for small shapes such as Please make workspace publication and consumption explicitly cache-safe, for example using the documented Test coverageThe added tests call
Please add at least one operator-level MatMul test that exercises the enabled dispatch and one fallback case. The current test inputs are also constant along K, which makes some K-indexing and split-boundary errors invisible. A random FP16 case checked against a higher-precision CPU reference would provide stronger coverage, especially for Counter contract and PR descriptionThe production path clears the counter before every launch, while the kernel also resets it and the test relies on that reset for a second launch. The header nevertheless says callers must provide a zeroed counter. Please choose and document one ownership/reset contract; if kernel-side reset is retained, use an explicitly atomic reset. The PR description also appears stale: it mentions a cached counter and mutex in Verdict: request changes. The overall integration is clean and the feature is safely default-off, but the cross-block publication protocol should be made CUDA-memory-model safe and the actual MatMul dispatch should be covered before merge. |
Decode-time projections such as the MoE router [2048, 256], the linear
attention in_proj gates [2048, 32] and the shared-expert gate [2048, 1]
have far too little N to fill a cuBLAS tile kernel: cuBLAS picks a 1x1
CTA grid and spends ~5 us reading a few hundred KiB of weights.
Add matmul_small_n_gemv.{h,cu}: grid (ceil(n/32), split_k), each block
accumulates one K-slice into an fp32 workspace, and the last block to
finish a column tile reduces the partials in slice order. Ordering the
reduction by slice index keeps the result deterministic, and the arrival
counter is reset by the reducing block so it only has to be cleared once
at allocation time rather than per launch.
The path is gated behind ORT_ENABLE_SMALL_N_GEMV and is OFF by default,
because in-model it is currently SLOWER than cuBLAS: 10.6 us vs 5.0 us
per call on H200 for the shapes above. A standalone microbenchmark had
suggested 2.1 us, but 2000 back-to-back launches keep the 131 KiB weight
resident in L2, which never happens in the real model. Landing it off
keeps the measurement and the kernel around for a future revisit; output
is token-identical to the cuBLAS path when enabled.
Make split-K workspace publication cache-safe and keep counter initialization caller-owned. Cover enabled operator dispatch, fallback, and varied K-axis inputs.
1128b51 to
e268362
Compare
|
Addressed the blocking feedback in e268362:
The four touched production/test translation units compile with the CUDA 13.0 Release build, and scoped lintrunner checks pass. A full |
…1478) ## Description Decode-time projections in hybrid / MoE LLMs are GEMVs with an `N` that is too small to fill a cuBLAS tile kernel. On the Qwen3.6-35B-A3B NVFP4 decode loop, router, linear-attention, and shared-expert gate projections repeatedly launch with `M <= 8`, `N <= 1024`, and `K >= 128`. This PR adds an experimental split-K GEMV kernel for that corner of the shape space. The path remains off by default because real-model measurements are currently slower than cuBLAS. ## Summary of Changes ### Kernel and dispatch - Adds a row-major FP16 split-K GEMV for `M <= 8`, `N <= 1024`, and `K >= 128`. - Gates dispatch with `ORT_ENABLE_SMALL_N_GEMV=1`; all ineligible layouts, shapes, transpose modes, and alpha values fall through to cuBLAS. - Reads the opt-in setting once when each MatMul kernel instance is created, keeping environment parsing out of `Compute()`. - Compares leading dimensions in `int64_t` so large shape values are not narrowed before eligibility checks. ### Cross-block reduction - Uses a grid of `(ceil(N / 32), split_k)` so K slices can run across multiple SMs. - Publishes FP32 partials through volatile global workspace accesses before `__threadfence()` and the completion atomic, following CUDA's last-block reduction visibility pattern. - Reduces partials in fixed slice order for deterministic output. - Uses per-invocation scratch for both workspace and completion counters. The caller clears counters before every launch; the kernel does not reset them. ### Tests - Direct-kernel tests cover every `M` specialization, `N = 1`, `N = 32`, `N = 1024`, uneven K splits, and repeated launches with explicit counter initialization. - Inputs vary along K and results are checked against an FP32 CPU reference. - Operator-level MatMul tests enable the feature and cover eligible `M = 8, N = 1` and `M = 8, N = 1024` dispatches plus an ineligible `K = 127` cuBLAS fallback. ## Current Status The path is default-off because it is currently slower than cuBLAS in the real model. On H200 (SM90) for the target decode shapes: | | us / call | us / decode step | |---|---:|---:| | cuBLAS | 5.0 | 671 | | this kernel | 10.6 | 1491 | A standalone microbenchmark had suggested approximately 2.1 us/call after subtracting the back-to-back launch floor. That result kept the small weight matrix resident in H200's L2, while the model reads cold weights. Keeping the implementation behind an opt-in flag preserves it for follow-up work that addresses that cold-read cost. ## Validation - Compiled `matmul.cc`, `matmul_small_n_gemv.cu`, and both small-N GEMV test translation units with the CUDA 13.0 Release build. - Ran scoped `lintrunner` checks for all changed source and test files. - Full `onnxruntime_test_all` linking is currently blocked by an unrelated `-Werror=unused-variable` in `contrib_ops/cuda/bert/paged_attention.cc` on the rebased main branch. ## Checklist - [x] No behavior change by default - [x] Cache-safe cross-block workspace publication - [x] Deterministic fixed-order reduction - [x] Operator-level enabled dispatch and fallback coverage - [ ] Enabled by default; blocked on closing the real-model performance gap against cuBLAS --------- Co-authored-by: GitHub Copilot <copilot@example.com>
…peed up MatMulNBits GEMV (#32876) ## Description Speculative decoding verifies a small batch of rows per step (8 for DFlash2 width 7), which lands two CUDA decode paths on poor kernels. Which kernel is best for these small-M shapes depends on the GPU: on SM 8.x cuBLAS picks a serial split-K kernel that is 10x+ slower than the small-N GEMV, on SM 12.0 the GEMV only helps at M = 2..8, and on SM 9.0 cuBLAS is as fast or faster. A fixed architecture policy regressed some of these (e.g. RTX 3060 at M = 9..64, N = 1024, up to 4.1x slower), so this PR selects the kernel per shape and per device by measuring it. 1. **GEMM auto-tuning for fp16/bf16 `MatMul` (opt-in).** With session config `ep.cuda.enable_gemm_auto_tune=1` (or env `ORT_CUDA_GEMM_AUTO_TUNE=1`), the first run of each eligible shape times cuBLAS and the small-N GEMV on the current device and caches the faster one for the process. When disabled (default), MatMul uses cuBLAS as on `main`. 2. **bf16 small-N GEMV.** The GEMV is templated on fp16/bf16 (fp32 accumulation, deterministic split-K). 3. **MatMulNBits fpA_intB GEMV at M >= 4.** The M x StepK activation tile plus the CtaM x CtaN accumulators push the 4/8-bit kernel to about 250 registers; halving CtaN frees registers and doubles the block count. The existing fpA_intB tactic profiler still chooses GEMV vs CUTLASS per M bucket. ## Summary of Changes ### GEMM auto-tuner | File | Change | |------|--------| | `include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h` | New key `ep.cuda.enable_gemm_auto_tune` (`kOrtSessionOptionsCudaEnableGemmAutoTune`). | | `onnxruntime/core/providers/cuda/math/gemm_auto_tuner.{h,cc}`, `gemm_auto_tuner_impl.cu` | Dispatch policy resolution, candidate timing, and a process-wide cache keyed by (device UUID, dtype, M, N, K, operand-alignment class). | | `onnxruntime/core/providers/cuda/math/matmul.{h,cc}` | `MatMul<MLFloat16/BFloat16>` resolves the policy at construction and, for eligible single GEMMs, looks up / tunes the kernel. Ineligible shapes (M > 64, N > 1024, K < 128, transposes, alpha != 1, batched) go straight to cuBLAS with no overhead. | Dispatch policy, in precedence order: | Setting | Behavior | |---|---| | `ORT_ENABLE_SMALL_N_GEMV=1` / `0` | Forces the small-N GEMV (when eligible) / cuBLAS; bypasses tuning. Kept for A/B testing and back-compat with the opt-in from #31478. | | session config `ep.cuda.enable_gemm_auto_tune` | `1` auto-tunes, `0` uses cuBLAS. | | env `ORT_CUDA_GEMM_AUTO_TUNE` | Fallback when the session config is unset. | | none | cuBLAS (same as `main`). | Tuning details: - **CUDA graph safe.** Tuning synchronizes the stream, so it is skipped while the stream is capturing; an untuned shape then uses cuBLAS and is not cached. ORT's warm-up runs before capture do the tuning. - **Representative timing.** Each timed run starts from a decode-like L2 state: a read-only pass over a scratch buffer (2x L2, capped at 256 MiB) evicts the weights without leaving dirty lines to write back, then the activation is read back in, since in decode it was just produced by the previous op (with a cold activation the H200 decisions disagreed with CUDA graph replay). A GPU delay kernel is queued first so host launch gaps of multi-launch candidates are hidden, candidates are interleaved, the fixed event-pair overhead (measured with an empty region) is subtracted, and the median of 10 runs is used. - **Stable choice.** cuBLAS is replaced only if the candidate is at least 5% faster. The cache is process-global (first insertion wins, tuning serialized by a mutex), so every session in a process runs the same kernel for the same shape. Each decision is logged at VERBOSE. - Cost: about 3 ms once per eligible (M, N, K) shape, M <= 64. ### Small-N GEMV | File | Change | |------|--------| | `onnxruntime/core/providers/cuda/math/matmul_small_n_gemv.cu`, `.h` | Covers M <= 64 in ordered 8-row chunks. Adds `SmallNGemvVecSplitKKernel` for even N, K % 8 == 0 and 16-byte aligned A. Each lane owns two columns and eight consecutive K rows, so one 16-byte broadcast A load feeds 16 FMAs; warps reduce through shared memory, then a deterministic last-block split-K reduction runs over `st.cg`/`ld.cg` partials. Completion counters are cleared by a one-block kernel instead of `cudaMemsetAsync` (a memset node costs several microseconds inside a CUDA graph). Templated on fp16/bf16. Every eligible vectorized shape splits K at least twice, so the unreachable single-slice store path was removed. | ### MatMulNBits fpA_intB GEMV | File | Change | |------|--------| | `onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemv/dispatcher.h` | For M >= 4, 4/8-bit weights (StepK < 64) use `CtaNLargeM = CtaN / 2` (8 -> 4 without zero points). M = 1..3 and the 2-bit layout are unchanged. | ## Testing - CUDA internal tests (`onnxruntime_providers_cuda_ut`, H200): - `GemmAutoTunerTest.*`: policy precedence and parsing, candidate selection margin, first-insertion-wins cache, timing/selection with synthetic GPU-delay candidates, cached keys not re-timed, CUDA graph capture detection. - `MatMulSmallNGemvTest.*`: fp16 and bf16 for M = 1..8, 17, vectorized/scalar variants, column-tile boundaries, stale counters cleared by the launcher, vectorized-kernel selection. - `MatMulSmallNGemvOpTest.*`: fp16 and bf16 MatMul through the CUDA EP, both forced and auto-tuned, for eligible shapes (M = 1..64, N = 1..1024) and ineligible fallbacks (K = 127, M = 65, N = 1025). - `onnxruntime_provider_test --gtest_filter='*MatMul*:*Gemm*'` (1003 tests) passes with and without `ORT_CUDA_GEMM_AUTO_TUNE=1`. - Python `test_matmul_gemm_auto_tune_cuda_graph_replay`: fp16 MatMuls tuned during warm-up, then captured and replayed with changing inputs. - Compiled for sm_86 and sm_90 in both the in-tree and the plugin CUDA EP builds. - H200 (SM 9.0), CUDA graph replay of enough MatMul copies that weights stream from DRAM, per-MatMul us, K = 5120, best of two interleaved sessions per mode. Auto-tune picks the GEMV only where it wins, and is within 1% of the faster kernel everywhere except one near-tie inside the 5% margin (bf16 4x48). Decisions were identical across five separate processes. | dtype | M | N | cuBLAS | forced GEMV | auto-tune (choice) | |---|---:|---:|---:|---:|---| | fp16 | 1 | 48 | 6.25 | 6.23 | 6.23 (GEMV) | | fp16 | 2 | 48 | 6.95 | 6.21 | 6.20 (GEMV) | | fp16 | 4 | 48 | 6.94 | 6.72 | 6.72 (GEMV) | | fp16 | 8 | 48 | 7.00 | 7.57 | 6.99 (cuBLAS) | | fp16 | 16 | 48 | 7.31 | 13.03 | 7.31 (cuBLAS) | | fp16 | 8 | 1024 | 8.17 | 13.60 | 8.19 (cuBLAS) | | fp16 | 16 | 1024 | 8.33 | 22.62 | 8.29 (cuBLAS) | | bf16 | 2 | 48 | 6.95 | 6.31 | 6.31 (GEMV) | | bf16 | 4 | 48 | 6.94 | 6.72 | 6.96 (cuBLAS) | | bf16 | 8 | 1024 | 8.28 | 13.89 | 8.17 (cuBLAS) | Tuning adds about 3 ms to the first run of a session per eligible shape. - Previous results for the kernels themselves (RTX 4090, sm_89, CUDA 13.3, nsys, weights streamed from DRAM), which auto-tuning now reaches when it selects the GEMV: | Shape (M x N x K) | cuBLAS / before | GEMV / after | |---|---:|---:| | fp16 MatMul 8 x 48 x 5120 | 78.6 us | 4.2 us | | fp16 MatMul 8 x 1024 x 5120 | 8.4 us | 7.3 us | | MatMulNBits GEMV 8 x 34816 x 5120 | 117.1 us | 112.0 us | | MatMulNBits GEMV 8 x 10240 x 5120 | 40.4 us | 35.7 us | | MatMulNBits GEMV 8 x 5120 x 17408 | 60.0 us | 57.3 us | | MatMulNBits GEMV 8 x 1024 x 5120 | 10.4 us | 8.9 us | End-to-end on RTX 4090 (ORT GenAI, CUDA plugin EP, Qwen3.8-27B INT4 + DFlash2 width 7, greedy, 1,024 tokens) the GEMV + MatMulNBits changes gave x1.28-x1.31 decode throughput (8/8 prompts per context) with unchanged TTFT and MMLU-Pro (88/100). With this revision, set `ep.cuda.enable_gemm_auto_tune=1` in the model's session options to get the MatMul part of that gain. ## Motivation and Context - The small-N GEMV (#31478, extended to 64 rows in #32289) was opt-in because on SM 9.0 cuBLAS was faster in the real model, while on SM 8.x cuBLAS is 10x+ slower at M = 2..8. An architecture default (the first revision of this PR) regressed RTX 3060 above M = 8 (4.1x at M = 64, N = 1024) and was neutral on RTX 5060 Ti; per-shape measurement on the actual device avoids hand-maintained thresholds. - A dedicated tuner is used instead of TunableOp: TunableOp is not built into the CUDA plugin EP, has no CUDA graph capture guard, times with a hot L2, sleeps 50 ms per tuned shape, and is enabled EP-wide. - Candidates are a table, so follow-ups can add kernels (e.g. cuBLASLt heuristic algorithms, CUTLASS, TRT-LLM `cudaCoreGemm`/tinygemm2, FlashInfer TGV on Blackwell) and a persistent TSV tuning cache like #29588. - The follow-up #32884 removes the remaining memset nodes from the decode graph (fpA_intB split-K, paged XQA, causal conv) and makes the fpA_intB tactic profiler prefer the GEMV on near ties at small M. ## Checklist - [x] Tests added/updated - [x] No breaking changes. Without the new option MatMul uses cuBLAS as on `main`; `ORT_ENABLE_SMALL_N_GEMV` keeps its meaning. - [x] Documentation updated (session option key comment) --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Description
Decode-time projections in hybrid / MoE LLMs are GEMVs with an
Nthat is too small to fill a cuBLAS tile kernel. On the Qwen3.6-35B-A3B NVFP4 decode loop, router, linear-attention, and shared-expert gate projections repeatedly launch withM <= 8,N <= 1024, andK >= 128.This PR adds an experimental split-K GEMV kernel for that corner of the shape space. The path remains off by default because real-model measurements are currently slower than cuBLAS.
Summary of Changes
Kernel and dispatch
M <= 8,N <= 1024, andK >= 128.ORT_ENABLE_SMALL_N_GEMV=1; all ineligible layouts, shapes, transpose modes, and alpha values fall through to cuBLAS.Compute().int64_tso large shape values are not narrowed before eligibility checks.Cross-block reduction
(ceil(N / 32), split_k)so K slices can run across multiple SMs.__threadfence()and the completion atomic, following CUDA's last-block reduction visibility pattern.Tests
Mspecialization,N = 1,N = 32,N = 1024, uneven K splits, and repeated launches with explicit counter initialization.M = 8, N = 1andM = 8, N = 1024dispatches plus an ineligibleK = 127cuBLAS fallback.Current Status
The path is default-off because it is currently slower than cuBLAS in the real model. On H200 (SM90) for the target decode shapes:
A standalone microbenchmark had suggested approximately 2.1 us/call after subtracting the back-to-back launch floor. That result kept the small weight matrix resident in H200's L2, while the model reads cold weights. Keeping the implementation behind an opt-in flag preserves it for follow-up work that addresses that cold-read cost.
Validation
matmul.cc,matmul_small_n_gemv.cu, and both small-N GEMV test translation units with the CUDA 13.0 Release build.lintrunnerchecks for all changed source and test files.onnxruntime_test_alllinking is currently blocked by an unrelated-Werror=unused-variableincontrib_ops/cuda/bert/paged_attention.ccon the rebased main branch.Checklist