Skip to content

[CUDA] Vectorize the NVFP4 weight dequantization for prefill - #32128

Merged
Tianlei Wu (tianleiwu) merged 2 commits into
microsoft:mainfrom
tianleiwu:tlwu/fp4_prefill_dequant_vec
Aug 18, 2026
Merged

Tianlei Wu (tianleiwu) merged 2 commits into
microsoft:mainfrom
tianleiwu:tlwu/fp4_prefill_dequant_vec

Conversation

@tianleiwu

Copy link
Copy Markdown
Contributor

Description

MatMulBlockQuantizedFp4Weight falls 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 DequantizeNvFp4Kernel gave each thread one packed byte (2 FP4 codes):

defect cost
1-byte load + two separate 2-byte stores no vectorization
idx / half_k and k0 / block_size two integer divisions per thread
*weight_scale_2 read per thread a global load per thread
__nv_cvt_fp4x2_to_halfraw2() software-emulated on pre-Blackwell: the SASS has branches and a subnormal normalization loop

DequantizeNvFp4Vec8Kernel replaces it when K % 8 == 0 and block_size is 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 from blockIdx.y, weight_scale_2 is hoisted into a register, and the scale index advances incrementally instead of by per-element division. Codes are decoded with the branch-free Fp4Cvt prmt lookup already in this file (added for the decode GEMV in #31155).

The scalar kernel is kept for odd block_size or K % 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 uint4 stores) 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 one uint4 store 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 Fp4Cvt reproduces the intrinsic's bit pattern exactly, so no output should change. Checked three ways:

  • Fp4Cvt vs __nv_cvt_fp4x2_to_halfraw2() brute-forced over all 256 packed byte values — 0 mismatches (including -0.0 for code 0x8).
  • Op-level SHA-256 over the full output for 8 shapes (both dtypes, block_size 16/32, K % 32 == 0, K % 32 == 16, K % 8 == 4, N/K from 128 to 5120) — identical before and after.
  • Qwen3.8-27B NVFP4, 8K prompt / 128 generated / MTP N=3 on 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):

N K scalar vectorized speedup
4096 4096 60.7 us 15.5 us 3.93x
6144 2048 46.1 us 12.1 us 3.81x
2048 6144 46.2 us 11.9 us 3.88x

Tests. Four new cases in matmul_block_scaled_fp4_test.cc, all with M > 8 so the decode GEMV is skipped and the dequant actually runs. Existing FP4 tests all used M <= 8, so the prefill path had no coverage at prefill shapes. Each case was confirmed under nsys to reach the intended kernel:

test kernel reached
PrefillDequantVectorizedFp16 DequantizeNvFp4Vec8Kernel<__half>
PrefillDequantVectorizedBiasBf16 DequantizeNvFp4Vec8Kernel<__nv_bfloat16>
PrefillDequantOddBlockSizeFp16 DequantizeNvFp4Kernel<__half>
PrefillDequantKNotMultipleOf8Bf16 DequantizeNvFp4Kernel<__nv_bfloat16>

PrefillDequantOddBlockSizeFp16 is the interesting one: with an odd block_size the 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 MatMulBlockQuantizedFp4WeightOpTest cases pass.

Motivation and Context

Measured on Qwen3.8-27B NVFP4 (168 MatMulBlockQuantizedFp4Weight nodes holding 14.97 G weights), 8K prompt, H200:

  • DequantizeNvFp4Kernel was 46.9% of all prefill GPU time (1750 ms of 3732 ms) — more than every cuBLAS GEMM in the model combined.
  • After this change it is 17.1% (409 ms), and total prefill GPU time drops 3732 -> 2394 ms.
  • End-to-end TTFT for 8K/128/spec-3: 3876 -> 2618 ms (-32.5%), averaged over 3 interleaved runs per arm. Decode throughput and MTP acceptance are unchanged, as expected — decode uses the GEMV path and never reaches this kernel.

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.

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.
Copilot AI balanced review requested due to automatic review settings August 17, 2026 08:34

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.

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.

Comment thread onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc
Comment thread onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu Outdated
@hariharans29

Copy link
Copy Markdown
Member

Are these all for Qwen 3.6 / Qwen 3.8 / GPT-OSS / applicable to all ?

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.
@tianleiwu

Copy link
Copy Markdown
Contributor Author

Are these all for Qwen 3.6 / Qwen 3.8 / GPT-OSS / applicable to all ?

For Qwen and DeepSeek. The GPT-OSS has MXFP4 in QMoE, which does not use this kernel.

@tianleiwu
Tianlei Wu (tianleiwu) merged commit 2d211f8 into microsoft:main Aug 18, 2026
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.
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.

4 participants