[CUDA] Vectorize the NVFP4 weight dequantization for prefill - #32128
Merged
Tianlei Wu (tianleiwu) merged 2 commits intoAug 18, 2026
Merged
Tianlei Wu (tianleiwu) merged 2 commits into
Tianlei Wu (tianleiwu) merged 2 commits into
Conversation
The prefill fallback of MatMulBlockQuantizedFp4Weight expands the packed weight into an [N, K] scratch buffer before calling cuBLAS. That expansion ran one thread per packed byte: a 1-byte load, two 2-byte stores, two integer divisions and a global load of weight_scale_2 per thread. DequantizeNvFp4Vec8Kernel gives each thread exactly one 8-element K chunk of one row, so a warp issues one contiguous 128-byte load and one contiguous 512-byte store. The row index comes from blockIdx.y instead of a 64-bit division, weight_scale_2 is hoisted into a register, and the scale index is advanced incrementally. Codes are decoded with the existing branch-free Fp4Cvt prmt lookup (added for the decode GEMV in microsoft#31155) rather than __nv_cvt_fp4x2_to_halfraw2(), which is emulated in software on SM90 with branches and a subnormal normalization loop. Fp4Cvt reproduces the intrinsic's bit pattern exactly, so the dequantized weight is bitwise identical to the scalar kernel. The scalar kernel is kept for odd block_size or K % 8 != 0.
Contributor
There was a problem hiding this comment.
Pull request overview
Optimizes CUDA NVFP4 prefill by vectorizing weight dequantization before cuBLAS.
Changes:
- Adds an 8-element-per-thread dequantization kernel and dispatch.
- Adds vectorized and scalar-fallback tests.
- Documents the optimized path and benchmarks.
Reviewed changes
Copilot reviewed 3 out of 3 changed files in this pull request and generated 2 comments.
| File | Description |
|---|---|
matmul_block_scaled_fp4.cu |
Implements and dispatches vectorized dequantization. |
matmul_block_scaled_fp4_test.cc |
Adds prefill-path tests. |
matmul_block_scaled_fp4.md |
Documents behavior and performance. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
This was referenced Aug 17, 2026
Baiju Meswani (baijumeswani)
approved these changes
Aug 18, 2026
Member
|
Are these all for Qwen 3.6 / Qwen 3.8 / GPT-OSS / applicable to all ? |
Hariharan Seshadri (hariharans29)
approved these changes
Aug 18, 2026
Tianlei Wu (tianleiwu)
added a commit
that referenced
this pull request
Aug 18, 2026
### Description
The `MatMulBlockQuantizedFp8Weight` fallback path (`M > 8`, i.e. every
prefill matmul) expands the packed weight into an `[N, K]` scratch
buffer of the activation type and then calls cuBLAS. That scratch is
**2.37 GiB** for a 248320 x 5120 LM head, and it stays live for the
duration of the GEMM even though the GEMM reads it exactly once.
This caps the scratch and runs the dequantize + GEMM pair over N tiles
that fit inside the cap:
```
for n_offset in 0, tile_rows, 2 * tile_rows, ...:
dequantize B[n_offset : n_offset + rows, :] into the scratch
cublasGemmHelper(...) writing Y + n_offset
```
The row-major `[M, N]` output is column-major `[N, M]` to cuBLAS, so an
N tile is a plain row offset into `Y` — there is no extra copy, no
split-K reduction, and no change to the arithmetic performed per output
element.
`ORT_FP8_DEQUANT_SCRATCH_MIB` sets the cap in MiB (default 256). Shapes
small enough to fit the cap take a single tile and are completely
unaffected.
### Verification
**Memory and speed** on Qwen3.8-27B at an 8K prompt (H200). The cap was
chosen by sweeping it:
| cap | peak memory | TTFT |
|---|---|---|
| untiled (before) | 32609 MiB | baseline |
| 1 GiB | -3100 MiB | +5.2% |
| **256 MiB (default)** | **28503 MiB (-4106)** | **+1.1%** |
| 128 MiB | -4106 MiB | +13.9% |
**Numerics.** Per-element arithmetic is unchanged, but tiling changes
the N extent handed to cuBLAS, so the library is free to select a
different kernel for a tile than for the whole matrix. Measured tiled vs
untiled, FP16, `K = 5120`, `block_size = 128`:
| M | N | tiles | max abs diff | differing elements |
|---:|---:|---:|---:|---:|
| 16 | 6144 | 1 | 0 | 0 |
| 1024 | 6144 | 1 | 0 | 0 |
| 16 | 32768 | 2 | 0.125 | 28065 |
| 1024 | 32768 | 2 | 0.125 | 2396209 |
| 64 | 248320 | 10 | 0 | 0 |
`0.125` is exactly one FP16 ULP at the magnitude of these outputs
(values around 200, where FP16 spacing is `2^-3`). So where the split
changes cuBLAS's kernel choice the result moves by at most one last-bit
rounding step — the same order of change as cuBLAS picking a different
tactic for any other reason. Single-tile shapes are bitwise identical.
**Tests.** `WeightDequantScratchTilingFp16` sets
`ORT_FP8_DEQUANT_SCRATCH_MIB=1` so a small shape (`N = 769`, `K = 4096`)
still splits into six tiles, with a distinct scale on every N row so
that a tile reading the wrong weight/scale offset or writing the wrong
output column changes `Y`. Confirmed under nsys that the test issues 3
dequant launches while a single-tile control issues 1. All 11
`MatMulBlockQuantizedFp8WeightOpTest` cases pass.
### Motivation and Context
On Qwen3.8-27B the untiled scratch was the single largest transient
allocation in the model and pushed peak memory ~4.1 GB above what the
weights and KV cache actually need. Recovering it makes room for longer
contexts or a larger batch at the same memory budget, and costs ~1% of
TTFT.
This is independent of #32128 (which speeds up the NVFP4 dequantization
kernel); the two touch different operators and can land in either order.
Contributor
Author
For Qwen and DeepSeek. The GPT-OSS has MXFP4 in QMoE, which does not use this kernel. |
Tianlei Wu (tianleiwu)
merged commit Aug 18, 2026
2d211f8
into
microsoft:main
90 of 91 checks passed
Tianlei Wu (tianleiwu)
added a commit
that referenced
this pull request
Aug 29, 2026
## Summary - Tune tensor-core NVFP4 GEMV dispatch to require multiple waves before selecting wide column tiles. - Preserve a larger K split for long reductions while avoiding excessive K splitting on Qwen gate/up shapes. - Add shape-scoped benchmark overrides, boundary tests, and CUDA contrib documentation. ## Motivation On H200, the previous dispatch could select a configuration with too few blocks for Qwen MTP shapes such as N=17408, K=5120. The updated policy keeps KSplit=8 for the longer K=8192 reduction while using KSplit=2 for K=5120, and requires sufficient grid waves before selecting wide column tiles. This PR is the tiling follow-up to #32128 and contains no duplicate vectorized NVFP4 dequantization changes. It can be rebased/stacked on #32128 after that PR merges. ## Validation - CUDA 13.0 / SM90 build passed. - `git diff --check` passed. - Added Qwen boundary coverage for the selected tensor-core tilings. - The local build configuration had provider unit tests disabled, so the new gtest was not executed locally.
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
MatMulBlockQuantizedFp4Weightfalls back to "expand the packed weight into an[N, K]scratch buffer, then cuBLAS" whenever the decode GEMV and the native SM120 path do not apply — i.e. for every prefill matmul on Hopper. That expansion was the single most expensive kernel in NVFP4 prefill.The old
DequantizeNvFp4Kernelgave each thread one packed byte (2 FP4 codes):idx / half_kandk0 / block_size*weight_scale_2read per thread__nv_cvt_fp4x2_to_halfraw2()DequantizeNvFp4Vec8Kernelreplaces it whenK % 8 == 0andblock_sizeis even. Each thread owns exactly one 8-element K chunk of one row, so a warp issues one contiguous 128-byte packed load and one contiguous 512-byte store. The row index comes fromblockIdx.y,weight_scale_2is hoisted into a register, and the scale index advances incrementally instead of by per-element division. Codes are decoded with the branch-freeFp4Cvtprmtlookup already in this file (added for the decode GEMV in #31155).The scalar kernel is kept for odd
block_sizeorK % 8 != 0.One design point worth recording: 8 elements per thread, not more. Widening the per-thread chunk to 32 elements (four back-to-back
uint4stores) makes every store instruction stride across lanes. Against a copy-only kernel with identical index math, that shape ceilings at 1.9 TB/s while oneuint4store per thread reaches 3.9 TB/s on H200 — a 2x difference that no amount of tile tuning recovers.Verification
Bitwise identical. Dequantization is elementwise and
Fp4Cvtreproduces the intrinsic's bit pattern exactly, so no output should change. Checked three ways:Fp4Cvtvs__nv_cvt_fp4x2_to_halfraw2()brute-forced over all 256 packed byte values — 0 mismatches (including-0.0for code0x8).block_size16/32,K % 32 == 0,K % 32 == 16,K % 8 == 4,N/Kfrom 128 to 5120) — identical before and after.N=3on H200: the generated-token SHA-256 is unchanged (095222fc5fdb…) across 3 runs per arm.Kernel time (H200,
M = 1024, BF16,block_size = 16, median over 50 iterations):Tests. Four new cases in
matmul_block_scaled_fp4_test.cc, all withM > 8so the decode GEMV is skipped and the dequant actually runs. Existing FP4 tests all usedM <= 8, so the prefill path had no coverage at prefill shapes. Each case was confirmed under nsys to reach the intended kernel:PrefillDequantVectorizedFp16DequantizeNvFp4Vec8Kernel<__half>PrefillDequantVectorizedBiasBf16DequantizeNvFp4Vec8Kernel<__nv_bfloat16>PrefillDequantOddBlockSizeFp16DequantizeNvFp4Kernel<__half>PrefillDequantKNotMultipleOf8Bf16DequantizeNvFp4Kernel<__nv_bfloat16>PrefillDequantOddBlockSizeFp16is the interesting one: with an oddblock_sizethe two nibbles of a packed byte can land in different scale blocks, which is exactly the assumption the vectorized kernel makes and therefore the reason it must be skipped.All 17
MatMulBlockQuantizedFp4WeightOpTestcases pass.Motivation and Context
Measured on Qwen3.8-27B NVFP4 (168
MatMulBlockQuantizedFp4Weightnodes holding 14.97 G weights), 8K prompt, H200:DequantizeNvFp4Kernelwas 46.9% of all prefill GPU time (1750 ms of 3732 ms) — more than every cuBLAS GEMM in the model combined.The kernel is now at the bandwidth bound for the work it does: one full dequantization pass over these weights moves 34.86 GiB, which is 12.5 ms at ~3 TB/s, and the measured cost is 12.4 ms per pass.