Skip to content

[CUDA] Auto-tune small-M decode MatMul (cuBLAS vs small-N GEMV) and speed up MatMulNBits GEMV - #32876

Merged
Tianlei Wu (tianleiwu) merged 3 commits into
mainfrom
tlwu/cuda-small-m-decode-gemv
Sep 29, 2026
Merged

Tianlei Wu (tianleiwu) merged 3 commits into
mainfrom
tlwu/cuda-small-m-decode-gemv

Conversation

@tianleiwu

@tianleiwu Tianlei Wu (tianleiwu) commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

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

Checklist

  • Tests added/updated
  • No breaking changes. Without the new option MatMul uses cuBLAS as on main; ORT_ENABLE_SMALL_N_GEMV keeps
    its meaning.
  • Documentation updated (session option key comment)

Enable the small-N fp16 split-K GEMV by default below SM 9.0 for up to
64 rows and add a vectorized kernel variant. cuBLAS picks a serial
split-K kernel for 2..8-row small-N fp16 shapes on SM 8.x that is more
than 10x slower.

Halve the fpA_intB GEMV CtaN for M >= 4 (4/8-bit) to cut register
pressure and double the block count for speculative-decode batches.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot AI balanced review requested due to automatic review settings September 28, 2026 18:19

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

The new default policy is untested, and the claimed single-split test actually selects two splits.

Review effort: Balanced
Findings: 2 Medium severity

Open (2)
What changed in this PR

Optimizes CUDA decode-time MatMul and MatMulNBits for small speculative-decoding batches.

Changes:

  • Enables small-N FP16 GEMV by default below SM 9.0.
  • Adds a vectorized split-K GEMV for eligible shapes.
  • Narrows MatMulNBits output tiles for larger row counts.
File Description
matmul.h Adds architecture-aware GEMV enablement.
matmul.cc Implements environment overrides and device defaults.
matmul_small_n_gemv.h Documents vectorized dispatch.
matmul_small_n_gemv.cu Adds vectorized kernels, dispatch, and workspace sizing.
dispatcher.h Narrows MatMulNBits tiles for M ≥ 4.
matmul_small_n_gemv_test.cc Adds vector-kernel test cases.

💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread onnxruntime/core/providers/cuda/math/matmul.cc Outdated
Comment thread onnxruntime/test/providers/cuda/test_cases/matmul_small_n_gemv_test.cc Outdated
…mset

Inside a CUDA graph a memset node costs several microseconds of dependency latency, while a kernel
node costs well under one, and decode issues one per small-N MatMul. The last block of each column
tile now also resets its counter, so one clear covers every row chunk of a launch.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Replace the SM < 9.0 default for the small-N GEMV with an opt-in per-shape,
per-device auto-tuner (session config ep.cuda.enable_gemm_auto_tune, env
ORT_CUDA_GEMM_AUTO_TUNE). Without it MatMul uses cuBLAS as on main;
ORT_ENABLE_SMALL_N_GEMV=1/0 still forces a kernel.

- gemm_auto_tuner: times candidates from a decode-like L2 state (weights
  evicted by a read-only pass, activation reloaded), hides host launch gaps
  behind a GPU delay kernel, subtracts the empty event-pair overhead, keeps
  cuBLAS unless another kernel is >= 5% faster, and caches the choice
  process-wide. Tuning is skipped while a CUDA graph is being captured.
- Small-N GEMV: templated on fp16/bf16; drop the unreachable single-split
  store path of the vectorized kernel.
- Tests: tuner unit tests, fp16/bf16 kernel and operator tests (forced and
  auto-tuned), and a Python CUDA graph replay test.
@tianleiwu Tianlei Wu (tianleiwu) changed the title [CUDA] Speed up small-M decode MatMul and MatMulNBits GEMV [CUDA] Auto-tune small-M decode MatMul (cuBLAS vs small-N GEMV) and speed up MatMulNBits GEMV Sep 29, 2026
Tianlei Wu (tianleiwu) added a commit to microsoft/onnxruntime-genai that referenced this pull request Sep 29, 2026
## Description

During CUDA speculative decoding (DFlash2 width 7, batch 1), the GPU sat
idle for about 2 ms of every ~27 ms step while the engine did host work
between the drafter and target graphs. On an RTX 4090 with Qwen3.8-27B,
three host-side costs accounted for most of that idle time: first-touch
page faults in the memory-mapped CPU embedding table, a pinned host
allocation and synchronizing free for every short-lived device view, and
a device round trip for every accepted draft token. This PR removes all
three. Token outputs are unchanged.

## Summary of Changes

| File | Change |
|------|--------|
| `src/ep/cuda/interface.cpp` | `PinnedHostPool` recycles the pinned
host mirrors (256 B to 1 MiB, power-of-two classes) of `GpuMemory`
views, so views no longer pin with `cudaHostAlloc` and release with
`cudaFreeHost`. `cudaFreeHost` also waits for the device. A mirror
released while a host-to-device copy from it may still be queued keeps
an event recorded behind that copy and is only handed out again after
the event completes. Larger mirrors keep the old path. |
| `src/ep/cuda/search_cuda.cpp` | `GreedySearch_Cuda::CommitToken` fast
path for a live batch-1 search and a non-EOS token. The host already
knows what `CheckForEOSAndPad` would do, so the token is written to the
next-token slot and the sequence with two async copies. This replaces
two kernels, a device-to-host copy and a stream synchronization per
accepted draft token. EOS, finished searches and batch > 1 keep the
existing path. |
| `src/models/cpu_embedding.cpp`, `.h` | `CpuEmbedding::Prefault` looks
every row up once, in chunks of 1,024 ids, when the session is created.
ORT maps CPU initializers stored as external data straight from the
file, so otherwise each decode lookup of a not-yet-touched row
page-faults. |
| `test/cpp/search_checkpoint_tests.cpp` |
`CommitTokenAppendsNonEosTokensUntilMaxLength` checks the fast path's
sequence contents, next token and max-length handling. |

## Testing

- `engine_unit_tests` with the CUDA plugin EP: all 806 tests pass,
including the new
`CudaSearchCheckpointTest.CommitTokenAppendsNonEosTokensUntilMaxLength`
and the existing EOS `CommitToken` tests.
- Token-exact in-process A/B: with a temporary switch between the old
and new commit and mirror paths, 16 prompts x 256 tokens at 0 and 8K
context generated identical tokens. The only difference was on one
prompt's first, cold-prefix-cache run, and two runs of the old path
differed the same way.
- CPU embedding lookups of 8 random rows over the 248,320 x 5120 fp16
table: about 1.1 ms before prefaulting and 12 us after. Prefaulting the
table takes about 1.8 s at load.
- Per-step host phases (nsys + NVTX, short prompt):

  | Phase | Before (us of GPU idle) | After |
  |---|---:|---:|
| CPU embedding, target + drafter (nested in the two input-preparation
rows) | 539 | 74 |
  | Accepted-token commit (`GenerateNextTokens`) | 328 | 229 |
  | DFlash2 input preparation | 329 | 103 |
  | Target decoder input preparation | 533 | 142 |
| Total GPU idle per step (after column includes the ORT graph changes)
| 2,828 | 1,641 |

- End-to-end, RTX 4090, Qwen3.8-27B INT4 + DFlash2, greedy, 16 prompts
(0 and 8K context) x 256 generated tokens, same ORT build, against two
baseline runs: the median paired decode step time went down 2.8-3.1% at
short context and 0.3-0.8% at 8K, and TTFT went down 4-6% because
prefill no longer takes embedding page faults. Together with the ORT
graph changes, the step went from 26.2 to 25.0 ms (short) and from 27.1
to 26.0 ms (8K) over 8 prompts x 1,024 tokens.

## Motivation and Context

The ORT-side decode graph overheads are handled separately in
microsoft/onnxruntime#32876 and microsoft/onnxruntime#32884.

## Checklist

- [x] Tests added/updated
- [x] No breaking changes
- [ ] Documentation updated (not applicable)

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Non-blocking suggestions: If the extra memory used for tuning cannot be allocated, use cuBLAS rather than fail a MatMul that could otherwise run. Also, the CUDA graph test enables tuning but may pick cuBLAS, so please add a test that forces the new GEMV path and checks graph replay. These can be follow-ups.

@tianleiwu
Tianlei Wu (tianleiwu) merged commit 567deb3 into main Sep 29, 2026
100 of 101 checks passed
@tianleiwu
Tianlei Wu (tianleiwu) deleted the tlwu/cuda-small-m-decode-gemv branch September 29, 2026 20:42
Tianlei Wu (tianleiwu) added a commit to microsoft/onnxruntime-genai that referenced this pull request Sep 30, 2026
…ode (#2643)

## Description

This follows #2642. It removes most of the remaining GPU idle time
between the drafter and target graphs in batch-1 CUDA speculative
decoding (DFlash2 width 7). On an RTX 4090 under Windows with
Qwen3.8-27B, nsys showed about 2 ms of idle GPU per ~26 ms step. Most of
it came from three places:

- **Copy-engine switches.** On WDDM, every switch between a copy-engine
operation (`cudaMemcpyAsync`, `cudaMemsetAsync`) and a kernel on the
same stream costs tens of microseconds of device idle time. The engine
issues dozens of tiny copies and memsets per step.
- **Synchronous state replay.** After a partially accepted verify, the
compact gated-delta-net/conv state replay (~0.37 ms) ran with the host
blocked on it.
- **Redundant host work.** The engine rebuilt ~400 fixed-state tensor
views every step, re-sampled a greedy token that verification had
already computed, and drained the stream in several places where it did
not need to.

## Summary of Changes

| Area | Change |
|------|--------|
| Small copies (`src/ep/cuda/interface.cpp`, `model_kernels.cu`,
`kernels.h`, `search_cuda.cpp`) | Small copies and memsets now run as
kernels:<ul><li>`GpuMemory` host-to-device uploads up to 4032 bytes use
`LaunchStoreBytes`, which passes the bytes as kernel parameters, so the
host mirror is free when the call returns.</li><li>Copies up to 256 KiB
in any direction use `LaunchCopyBytes`, which reads or writes the pinned
mirror through unified addressing.</li><li>Memsets up to 256 KiB use
`LaunchZeroBytes`.</li><li>The greedy argmax read-back uses a strided
gather kernel that writes the pinned buffer directly.</li><li>The search
transaction checkpoint copies, `ResetDone`, `RewindTo`, and
`CommitToken` token stores also use these kernels.</li></ul>Larger
transfers keep the copy engine. |
| GDN state replay (`src/engine/fixed_state_pool.cpp`,
`src/ep/cuda/interface.cpp`, `model_kernels.cu`) | `PrepareCommit` no
longer launches the compact replay and waits for it:<ul><li>The replay
descriptors are deferred to the pool's next operation (`Reserve`,
`CapturePrefixCheckpoint`, or `PrepareCommit`), which enqueues them
ahead of every later reader of the banks on the same stream. The device
runs the replay while the host prepares the next
step.</li><li>Discarding a prepared reservation drops its deferred
replay.</li><li>A launch failure marks the pool unhealthy.</li><li>A new
`ReplayGatedDeltaNetKernel` stages each head's decays, keys, and deltas
in shared memory and streams whole state rows as `float4`: about 300 us
instead of 320-525 us for the 151 MB of fp32 state.</li><li>Conv states
and ineligible geometries keep the generic kernel.</li></ul> |
| Fixed-state reservations (`src/engine/fixed_state_pool.cpp`) |
<ul><li>Tensor views and bindings are cached per layout (row count,
direct bank, first slot, capture mode) in `FixedStateReservationViews`,
which every reservation holds by `shared_ptr`.</li><li>On CUDA,
direct-binding reservations no longer drain the stream after the
capture-count upload.</li><li>Replay descriptors are built from raw
tensor pointers.</li></ul> |
| Greedy drafted rows (`src/engine/scheduled_requests.cpp`, `.h`) |
Drafted requests already run without logits processors
(`Request::DraftTokenValidationError`), so the correction/bonus token of
a greedy drafted request is the verification argmax of the row after the
accepted prefix.<ul><li>That token is committed with `CommitToken`
instead of sending the row through the batched sampler
again.</li><li>These requests take the external-sampling checkpoint, as
before.</li><li>They get a private next-token slot for the transaction,
so they never write into a batched-sampler slot that another request now
owns.</li></ul> |
| Graph runs (`src/engine/decoders/simple_decoder.cpp`,
`src/dflash2_drafter.cpp`) | Captured CUDA-graph replays of the target
run with `disable_synchronize_execution_providers=1`. The first host
read of their outputs (the sampled token ids) synchronizes the stream.
The drafter does the same when it reads drafts back. |
| CPU embedding (`src/models/cpu_embedding.cpp`, `src/smartptrs.h`) |
New `DeviceInterface::RecyclesHostMirrorsAfterUpload`; CUDA returns true
for mirrors up to 1 MiB. When it holds, each lookup takes a fresh pooled
pinned mirror instead of synchronizing the stream before it reuses one.
`kDeviceInterfaceVersion` goes to 7. |
| Docs | `paged_attention_engine.md` describes the deferred replay and
the greedy drafted-token commit. `cpu-embedding-engine.md` describes the
mirror handling. |
| Tests (`test/cpp/engine/fixed_state_pool_tests.cpp`) | New
tests:<ul><li>`CudaFixedStatePoolTest.GatedDeltaNetReplayMatchesHostRecurrence`
checks the fast and generic kernels bit-exactly against the host
recurrence.</li><li>`PrefixCheckpointAfterPartialAcceptanceSeesReplayedState`
covers a checkpoint captured right after a deferred
replay.</li><li>`DiscardedPartialAcceptanceLeavesStateUnchanged` covers
dropping the replay on discard.</li></ul> |

## Testing

- `engine_unit_tests` with the CUDA plugin EP: 809/809 pass, including
the three new tests.
- Microbenchmark (RTX 4090, WDDM): after a busy kernel, 6 small copies
interleaved with kernels added 274 us (host to device) and 343 us
(device to host) of wall time as `cudaMemcpyAsync`, and 9 us as copy
kernels.
- nsys + temporary NVTX, short prompt, 48 steps, host pinned to the
performance cores:

  | Per step | Before | After |
  |---|---:|---:|
  | Step span | 26.19 ms | 25.00 ms |
| GPU idle (incl. ~0.4 ms of in-graph launch gaps) | 1.98 ms | 1.16 ms |
  | Fixed-state reserve + commit host time | ~0.56 ms | ~0.07 ms |
| GDN replay on the critical path | 0.37 ms (host blocked) | mostly
overlapped |

- End-to-end A/B against the #2642 build (Qwen3.8-27B INT4 + DFlash2,
greedy, 8 prompts x 256 tokens, 2 interleaved runs each, minimum per
prompt): the median paired step time went down 1.4% at short context,
1.9% at 8K and 3.7% at 32K. The desktop was active during the runs, so
single runs vary by several percent.
- Greedy outputs: repeated runs of the #2642 build already differ on 2
of 8 prompts. The new build matched it on 6-7 of 8, which is within that
run-to-run variation.

## Motivation and Context

These changes build on #2642 (pinned mirror pool, `CommitToken` fast
path, CPU embedding prefault). The ORT-side graph changes are in
microsoft/onnxruntime#32876 and microsoft/onnxruntime#32884.

## Checklist

- [x] Tests added/updated
- [x] No breaking changes (`kDeviceInterfaceVersion` bumped for the new
virtual)
- [x] Documentation updated

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
Tianlei Wu (tianleiwu) added a commit that referenced this pull request Sep 30, 2026
… small M (#32884)

## Description

Speculative decoding on a hybrid (gated delta net + full attention)
decoder replays one CUDA graph per verify step, and that graph carried
up to about 130 memset nodes: serial split-K semaphores of the fpA_intB
CUTLASS GEMM, paged XQA semaphores, and the `VarlenCausalConvWithState`
`state_update` output. Inside a graph each memset node adds several
microseconds of dependency latency, while a small kernel adds well under
one. This PR clears those buffers with kernels and makes the weight-only
tactic profiler stop picking CUTLASS over the CUDA GEMV on near ties at
small M, which it does because it times L2-resident synthetic weights.

## Summary of Changes

### Replace in-graph memsets with kernels

| File | Change |
|------|--------|
|
`onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/device/gemm_universal_base_compat.h`
| `ClearSplitKSemaphoresKernel` replaces `cudaMemsetAsync` for the
serial split-K semaphore workspace (int-sized workspaces up to 4 MiB;
anything else keeps the memset). |
| `onnxruntime/contrib_ops/cuda/bert/xqa/xqa_paged_impl_gen.cuh`,
`xqa_paged_loader_impl.cuh` | Paged XQA semaphores are cleared with
`onnxruntime::cuda::Fill<int32_t>`. |
| `onnxruntime/contrib_ops/cuda/bert/varlen_causal_conv_with_state.cc` |
The `state_update` output is zeroed with `Fill<int32_t>` when its size
is a multiple of 4 bytes, and with `cudaMemsetAsync` otherwise. |

### fpA_intB tactic selection at small M

| File | Change |
|------|--------|
| `onnxruntime/contrib_ops/cuda/llm/gemm_profiler.h` | New virtual
`getSelectionTime(m, tactic, time)` hook. The default returns the
measured time unchanged, so other profilers are unaffected. |
| `onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemm_profiler.h`, `.cc` |
For M < 16 (the range where the CUDA GEMV is a candidate), a CUTLASS
tactic's time is multiplied by 1.1, so CUTLASS has to be more than 10%
faster to replace the GEMV. |

## Testing

- Synthetic graph on RTX 4090 (sm_89, CUDA 13.3, 200 tiny kernels): 1.0
us per kernel alone, 12.1 us with a memset node after each kernel, and
1.8 us with a one-block zeroing kernel after each kernel.
- The profiler's own M = 8 timings for the Qwen3.8-27B projections put
CUTLASS within about 1% of the GEMV on some shapes, for example q_proj
5120 -> 12288: 23.4 us CUTLASS vs 23.7 us GEMV. In the model, where
weights stream from DRAM, the same shapes run 45.7 us on CUTLASS and
43.2 us on the GEMV. Across the projections, the GEMV was 4-30% faster
in-model at M = 8.
- End-to-end on RTX 4090, ORT GenAI CUDA plugin EP, Qwen3.8-27B INT4 +
INT8 paged KV + DFlash2 width 7, greedy:
- Forcing the GEMV for every M < 16 bucket (a temporary experiment) cut
the mean decode step from 27.28 to 26.75 ms over 16 prompts, in two
interleaved repeats. With this PR the profiler picks the GEMV for every
M = 8 projection.
- nsys node-level traces of the target graph, per verify step: before,
the memset nodes and their neighbouring gaps cost about 725 us (small-N
GEMV counters 326 us, CUTLASS split-K semaphores 296 us, XQA 58 us,
causal conv 45 us), and the idle time between graph nodes was 723 us.
With this PR and the small-N counter change in #32876, no memset nodes
remain and the idle time between nodes is 306 us.
- Correctness through the plugin EP against a float64 reference,
MatMulNBits 4-bit block 32 with M in {4, 8, 12, 16, 24, 32, 48, 64} on
the split-K shapes 5120x6144, 1024x5120, 5120x17408 and 10240x5120:
worst relative error 4.3e-3, identical to the build without this change.
The trace confirms `ClearSplitKSemaphoresKernel` ran for every split-K
launch.
- MMLU-Pro (100 items, same protocol as #32876): 86/100 vs 88/100. The
kernel selection change alters rounding, so 46 of the 100 generations
diverge at some token. Two answers flipped from correct to incorrect,
and all outputs were well formed (no unextracted answers or
speculative-decoding failures).
- Suggested CI coverage: the MatMulNBits CUDA tests, `PagedAttention` /
XQA tests, and the `VarlenCausalConvWithState` tests.

## Motivation and Context

- The memsets sit in hot decode paths. Per verify step on the model
above: 64 CUTLASS split-K clears (when CUTLASS is selected), 16 XQA
clears, and 48 causal-conv clears.
- Timing the profiler with its synthetic weights evicted from L2 was
tried in #32876 and did not change end-to-end time, so this PR adds a
selection margin rather than changing how tactics are timed.
- Complements #32876, which makes the same change for the small-N fp16
GEMV counters.

## Checklist

- [ ] Tests added/updated (behavior-preserving; covered by existing op
tests)
- [x] No breaking changes
- [ ] Documentation updated (not applicable)

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants