[CUDA] Extend speculative decode GEMVs to 64 rows - #32289
Merged
Merged
Conversation
Contributor
There was a problem hiding this comment.
Pull request overview
Extends CUDA speculative-decoding GEMV paths to support up to 64 rows while retaining specialized fused kernels.
Changes:
- Processes standard small-N GEMV in 8-row chunks.
- Adds multi-tile and split-launch FP8/FP4 GEMV support.
- Expands dispatch-boundary tests through 64 rows.
Reviewed changes
Copilot reviewed 12 out of 12 changed files in this pull request and generated 3 comments.
Show a summary per file
| File | Description |
|---|---|
onnxruntime/test/providers/cuda/test_cases/matmul_small_n_gemv_op_test.cc |
Tests larger eligible row counts. |
onnxruntime/test/contrib_ops/matmul_block_scaled_fp8_test.cc |
Expands FP8 tile and fallback coverage. |
onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc |
Expands FP4 FP16/BF16 coverage. |
onnxruntime/core/providers/cuda/math/matmul.cc |
Delegates counter initialization to the launcher. |
onnxruntime/core/providers/cuda/math/matmul_small_n_gemv.h |
Documents chunked row processing. |
onnxruntime/core/providers/cuda/math/matmul_small_n_gemv.cu |
Implements 8-row chunked launches through M=64. |
onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp8.h |
Exposes FP8 GEMV row-limit selection. |
onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp8.cu |
Adds multi-tile and split-launch FP8 kernels. |
onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp8.cc |
Routes eligible larger shapes to FP8 GEMV. |
onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h |
Exposes FP4 GEMV row-limit selection. |
onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu |
Adds multi-tile and split-launch FP4 kernels. |
onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc |
Routes eligible larger shapes to FP4 GEMV. |
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
H200 measurements show 10-17% lower median latency at M=1, while M>=8 is slower than cuBLAS by up to 7.3x. Keep the direct launcher capable through M=64 for tests and future tuning.
H200 measurements avoid up to 39% FP4 and 15% FP8 regressions at larger M while retaining the 64-row path behind per-format overrides.
Tianlei Wu (tianleiwu)
requested review from
Baiju Meswani (baijumeswani),
Hariharan Seshadri (hariharans29) and
kunal-vaishnavi
August 29, 2026 20:15
Baiju Meswani (baijumeswani)
previously approved these changes
Aug 29, 2026
Baiju Meswani (baijumeswani)
approved these changes
Aug 30, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
Motivation and Context
Multi-token speculative verification produces matrix row counts beyond the original decode-only range. Falling back at those boundaries adds dequantization, workspace, and general GEMM overhead to a latency-sensitive path. This change keeps those shapes on the specialized kernels while preserving the existing fallback for ineligible dimensions.
Validation
MatMulSmallNGemvOpTest.DispatchesEligibleShapesWhenEnabledpassed and covers 8, 9, 33, and 64 rows.MatMulBlockQuantizedFp4WeightOpTest.GemvTensorCoreTilesFp16MatMulBlockQuantizedFp4WeightOpTest.GemvTensorCoreTilesBf16MatMulBlockQuantizedFp8WeightOpTest.GemvTensorCoreTilesFp16git diff --checkpassed for all 12 changed files.