From 0be9ba6d5dcb07830bf77323a7b76fab676a4a31 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Tue, 21 Jul 2026 16:48:13 -0700 Subject: [PATCH 01/11] [CUDA] Add MatMulBlockScaledFp4 contrib operator Add the com.microsoft MatMulBlockScaledFp4 CUDA contrib operator, a weight-only NVFP4 (E2M1) matmul that computes Y = A * dequant(B)^T (+bias) with packed 4-bit weights and per-block E4M3 scales plus a global fp32 scale. Includes an architecture-independent dequant + GEMM/GEMV path and a native SM120 CUTLASS NVFP4 path, the ONNX schema, kernel registration, build wiring, unit tests, docs and a profiling script. --- cmake/onnxruntime_cuda_source_filters.cmake | 3 +- cmake/onnxruntime_providers_cuda.cmake | 4 + cmake/onnxruntime_providers_cuda_plugin.cmake | 1 + .../cuda/matmul_block_scaled_fp4.md | 238 +++++++++++ .../matmul_block_scaled_fp4_experiments.md | 181 ++++++++ .../contrib_ops/cuda/cuda_contrib_kernels.cc | 2 + .../cuda/math/matmul_block_scaled_fp4.cc | 284 +++++++++++++ .../cuda/math/matmul_block_scaled_fp4.cu | 349 +++++++++++++++ .../cuda/math/matmul_block_scaled_fp4.h | 109 +++++ .../math/matmul_block_scaled_fp4_sm120.cu | 360 ++++++++++++++++ .../core/graph/contrib_ops/contrib_defs.cc | 57 +++ onnxruntime/core/graph/contrib_ops/ms_opset.h | 2 + .../matmul_block_scaled_fp4_test.cc | 207 +++++++++ .../profile_matmul_block_scaled.py | 398 ++++++++++++++++++ 14 files changed, 2194 insertions(+), 1 deletion(-) create mode 100644 docs/contrib_ops/cuda/matmul_block_scaled_fp4.md create mode 100644 docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md create mode 100644 onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc create mode 100644 onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu create mode 100644 onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h create mode 100644 onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu create mode 100644 onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc create mode 100644 onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py diff --git a/cmake/onnxruntime_cuda_source_filters.cmake b/cmake/onnxruntime_cuda_source_filters.cmake index 0b33675ea96f5..dc9a6b3ce433a 100644 --- a/cmake/onnxruntime_cuda_source_filters.cmake +++ b/cmake/onnxruntime_cuda_source_filters.cmake @@ -83,7 +83,8 @@ function(onnxruntime_extract_sm_specific_cuda_sources CU_SRC_LIST) set(_sm120_srcs) if("120" IN_LIST CMAKE_CUDA_ARCHITECTURES_ORIG) foreach(_src IN LISTS _list) - if(_src MATCHES "moe_gemm_tma_ws_sm120_.*\\.generated\\.cu$") + if(_src MATCHES "moe_gemm_tma_ws_sm120_.*\\.generated\\.cu$" OR + _src MATCHES "matmul_block_scaled_fp4_sm120\\.cu$") list(APPEND _sm120_srcs "${_src}") endif() endforeach() diff --git a/cmake/onnxruntime_providers_cuda.cmake b/cmake/onnxruntime_providers_cuda.cmake index c5fc0fd53ece1..1ab78a6aba243 100644 --- a/cmake/onnxruntime_providers_cuda.cmake +++ b/cmake/onnxruntime_providers_cuda.cmake @@ -586,6 +586,10 @@ CUDA_ARCHITECTURES "${_ort_sm120_cuda_architectures}" NVCC_THREADS "${onnxruntime_NVCC_THREADS}" SOURCES ${onnxruntime_cuda_sm120_tma_srcs}) + target_compile_definitions(onnxruntime_providers_cuda PRIVATE ORT_ENABLE_BLOCKQUANT_SM120) + if(TARGET onnxruntime_providers_cuda_obj) + target_compile_definitions(onnxruntime_providers_cuda_obj PRIVATE ORT_ENABLE_BLOCKQUANT_SM120) + endif() endif() endif() diff --git a/cmake/onnxruntime_providers_cuda_plugin.cmake b/cmake/onnxruntime_providers_cuda_plugin.cmake index bd126f92807b4..fb8750fa2c850 100644 --- a/cmake/onnxruntime_providers_cuda_plugin.cmake +++ b/cmake/onnxruntime_providers_cuda_plugin.cmake @@ -327,6 +327,7 @@ if(NOT onnxruntime_DISABLE_CONTRIB_OPS) NVCC_THREADS "${onnxruntime_plugin_nvcc_threads}" COMPILE_OPTIONS ${_cuda_plugin_shared_compile_options} SOURCES ${_cuda_plugin_sm120_tma_srcs}) + target_compile_definitions(onnxruntime_providers_cuda_plugin PRIVATE ORT_ENABLE_BLOCKQUANT_SM120) endif() endif() diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md new file mode 100644 index 0000000000000..14e2a03fcc9f1 --- /dev/null +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md @@ -0,0 +1,238 @@ +# MatMulBlockScaledFp4 - CUDA Operator Documentation + +This document describes the CUDA execution-provider implementation of +**MatMulBlockScaledFp4** (`com.microsoft::MatMulBlockScaledFp4`): its tensor +format, dispatch chain, native Blackwell path, prepacking behavior, and test / +benchmark workflow. + +MatMulBlockScaledFp4 computes `Y = A * dequant(B)^T (+ bias)` where `A` is +FP16 or BF16 and `B` is an `N x K` weight matrix stored as packed NVIDIA FP4 +E2M1 values with block-wise E4M3 scales. The default semantics are +weight-only FP4: activations stay FP16/BF16. An opt-in SM120 path quantizes +activations to NVFP4 internally and uses native block-scaled tensor cores. + +Source files: + +- [onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc](../../../onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc) - operator, validation, dispatch, and `PrePack`. +- [onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h](../../../onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h) - kernel class and CUDA launcher declarations. +- [onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu](../../../onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu) - dequantization, bias add, and decode GEMV kernels. +- [onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu](../../../onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu) - native SM120 NVFP4 x NVFP4 CUTLASS path. +- [onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py](../../../onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py) - opt-in accuracy and latency harness. + +--- + +## Table of Contents + +1. [Operator Schema](#1-operator-schema) +2. [Weight Format](#2-weight-format) +3. [Dispatch Chain](#3-dispatch-chain) +4. [Decode Path - Fused GEMV](#4-decode-path---fused-gemv) +5. [Default Path - Dequantize + cuBLAS](#5-default-path---dequantize--cublas) +6. [Native SM120 FP4 x FP4 Path](#6-native-sm120-fp4-x-fp4-path) +7. [PrePack](#7-prepack) +8. [Environment Variables](#8-environment-variables) +9. [Testing and Benchmarking](#9-testing-and-benchmarking) + +--- + +## 1. Operator Schema + +| Attribute | Meaning | +|-----------|---------| +| `K` | Input feature dimension: columns of `A` and logical columns of `B`. | +| `N` | Output feature dimension: rows of logical `B`. | +| `block_size` | Quantization group size along `K`. Current CUDA paths are optimized for `16`; default is `16`. | + +| Input | Index | Type | Notes | +|-------|-------|------|-------| +| `A` | 0 | FP16 or BF16 | Activation tensor with last dimension `K`. Leading dimensions are flattened into `M`. | +| `B` | 1 | UINT8 | Packed NVFP4 E2M1 weight, shape `[N, K / 2]`. Two FP4 values per byte, low nibble first. | +| `weight_scale` | 2 | UINT8 | Raw E4M3 per-block scales, shape `[N, ceil(K / block_size)]`. | +| `weight_scale_2` | 3 | FP32 scalar | Global weight scale. | +| `input_scale` | 4 | Optional FP32 scalar | Used only by the opt-in native SM120 FP4 x FP4 path. | +| `bias` | 5 | Optional FP16/BF16 | Bias of shape `[N]`, same type as `A`. | + +Output `Y` has the same leading dimensions as `A` and last dimension `N`. Its +type matches `A`. + +--- + +## 2. Weight Format + +`B` is a row-major logical `[N, K]` matrix packed to `[N, K / 2]` bytes. Each +byte contains two E2M1 values: + +- low nibble: even K element, +- high nibble: odd K element. + +`weight_scale[n, kb]` is a raw E4M3 byte for output row `n` and K block `kb`. +The dequantized value is: + +``` +B_dequant[n, k] = fp4_e2m1(B[n, k]) * e4m3(weight_scale[n, k / block_size]) * weight_scale_2 +``` + +`K` must be even because two FP4 values are packed per byte. The decode and +native SM120 paths additionally require `block_size == 16` and `K % 32 == 0`. + +--- + +## 3. Dispatch Chain + +`MatMulBlockScaledFp4::ComputeImpl` tries the cheapest applicable path first: + +```mermaid +flowchart TD + A[ComputeImpl] --> Z{empty output?} + Z -- yes --> R[return] + Z -- no --> G{M <= 8
block_size == 16
K % 32 == 0} + G -- yes --> GEMV[fused FP4 weight-only GEMV] --> R + G -- no --> N{native SM120 env enabled
SM120 device
block_size == 16
K % 32 == 0
N % 32 == 0} + N -- yes --> P[native NVFP4 x NVFP4 GEMM] --> BIAS[optional bias add] --> R + N -- no --> DQ[dequantize B to FP16/BF16 scratch] --> CUBLAS[cuBLAS GEMM] --> BIAS2[optional bias add] --> R +``` + +The decode GEMV path intentionally has priority over native SM120 GEMM. For +small `M`, the warp-per-column GEMV is memory-bound and avoids activation +quantization, CUTLASS setup, and underutilized tensor-core GEMM work. + +--- + +## 4. Decode Path - Fused GEMV + +`LaunchMatMulBlockScaledFp4Gemv` is used when: + +- `0 < M <= 8`, +- `block_size == 16`, +- `K % 32 == 0`. + +Each warp computes one output element `Y[row, col]`. A lane consumes 32 K +elements per iteration, which is exactly two 16-element scale blocks. The kernel +loads: + +- 16 packed FP4 bytes from one row of `B`, +- 32 FP16/BF16 activation values from `A`, +- two contiguous E4M3 scale bytes from `weight_scale[col, :]`. + +The per-block scales are folded into the partial sums and `weight_scale_2` is +applied once after the warp reduction. Optional bias is fused in lane 0. + +This kernel reads the original unswizzled `[N, K / 16]` scale layout. Experiments +with the native SM120 swizzled scale layout for GEMV were slower; see +[matmul_block_scaled_fp4_experiments.md](matmul_block_scaled_fp4_experiments.md). + +--- + +## 5. Default Path - Dequantize + cuBLAS + +When decode GEMV and native SM120 GEMM do not apply, the operator uses a +portable weight-only fallback: + +1. `LaunchDequantizeNvFp4` expands `B` into a scratch `[N, K]` buffer of the + activation type (FP16 or BF16). +2. cuBLAS computes `Y = A * B_dequant^T`. +3. `LaunchAddBiasNvFp4` adds optional bias. + +This path keeps full-precision activations and runs on CUDA devices with NVFP4 +conversion intrinsic support in the configured CUDA toolkit. It is the default +prefill path when the SM120 native environment variable is not enabled. + +--- + +## 6. Native SM120 FP4 x FP4 Path + +The native Blackwell path is compiled when the build defines +`ORT_ENABLE_BLOCKQUANT_SM120` and is enabled at runtime with: + +```bash +ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 +``` + +Runtime guards: + +- device compute capability is SM120 (`sm_ >= 120 && sm_ < 130`), +- `block_size == 16`, +- `K % 32 == 0`, +- `N % 32 == 0`, +- `M > 8` because decode GEMV has priority. + +The native path performs three steps: + +1. Quantize activation `A` to packed NVFP4 E2M1 with per-16-block E4M3 scales. +2. Provide `B` scales in the SM120 block-scaled swizzled layout required by + CUTLASS. If `PrePack` cached this layout, the cached buffer is reused; + otherwise it is repacked into scratch for this run. +3. Run CUTLASS block-scaled NVFP4 x NVFP4 GEMM and optionally add bias. + +Accuracy note: this path changes internal arithmetic from weight-only FP4 to +activation-and-weight FP4. The profiling harness therefore compares native SM120 +results against an activation-quantized FP4 reference when the env var and shape +select this path. + +--- + +## 7. PrePack + +`PrePack` handles input index `2` (`weight_scale`) only for the eligible native +SM120 path. It converts the original `[N, K / 16]` E4M3 scale tensor into the +SM120 swizzled scale layout once and stores it in `b_scale_prepacked_`. + +`is_packed` deliberately remains `false`: the original `weight_scale` input must +stay available because the decode GEMV and default dequant+cuBLAS paths still +consume the unswizzled layout. + +If `weight_scale` is not an initializer, or the native SM120 path is not enabled +or supported, the operator falls back to per-run scratch repacking for native +GEMM and the original scale tensor for the other paths. + +--- + +## 8. Environment Variables + +| Variable | Default | Meaning | +|----------|---------|---------| +| `ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120` | `0` | Enables the opt-in native SM120 NVFP4 x NVFP4 GEMM path when the shape and device guards pass. | + +The default remains the existing weight-only semantics: decode GEMV for small +`M`, otherwise dequantize `B` and call cuBLAS. + +--- + +## 9. Testing and Benchmarking + +Focused C++ tests: + +```bash +CUDA_VISIBLE_DEVICES=0 build/cu130/Release/onnxruntime_provider_test \ + --gtest_filter='MatMulBlockScaledFp4OpTest.*' +``` + +Python harness examples: + +```bash +# Decode GEMV +cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ + python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + --op fp4 --activation-dtype fp16 --m 1 --n 11008 --k 4096 --warmup 100 --repeat 500 + +# Default prefill: dequantize + cuBLAS +cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ + python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + --op fp4 --activation-dtype fp16 --m 16 --n 11008 --k 4096 --warmup 50 --repeat 200 + +# Native SM120 prefill +cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + --op fp4 --activation-dtype fp16 --m 16 --n 11008 --k 4096 --warmup 50 --repeat 200 +``` + +After rebuilding `libonnxruntime_providers_cuda.so`, sync the provider into the +Python load locations before Python benchmarks: + +```bash +cp build/cu130/Release/libonnxruntime_providers_cuda.so \ + build/cu130/Release/onnxruntime/capi/libonnxruntime_providers_cuda.so +cp build/cu130/Release/libonnxruntime_providers_cuda.so \ + build/cu130/Release/build/lib/onnxruntime/capi/libonnxruntime_providers_cuda.so +``` diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md new file mode 100644 index 0000000000000..e940f8ba91dd0 --- /dev/null +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md @@ -0,0 +1,181 @@ +# MatMulBlockScaledFp4 - CUDA Experiments + +This document records CUDA experiments for +**MatMulBlockScaledFp4** (`com.microsoft::MatMulBlockScaledFp4`) that are useful +for future performance work but are not part of the final dispatch chain. + +Related documentation: + +- [matmul_block_scaled_fp4.md](matmul_block_scaled_fp4.md) - operator behavior and current dispatch chain. +- [matmul_nbits_small_m_experiments.md](matmul_nbits_small_m_experiments.md) - similar small-M GEMV experiment notes for MatMulNBits. + +--- + +## Table of Contents + +1. [Native SM120 FP4 x FP4 GEMM](#1-native-sm120-fp4-x-fp4-gemm) +2. [Prepacking SM120 Swizzled B Scales](#2-prepacking-sm120-swizzled-b-scales) +3. [Rejected Experiment - Reuse Swizzled B Scales for Decode GEMV](#3-rejected-experiment---reuse-swizzled-b-scales-for-decode-gemv) +4. [Benchmark Commands](#4-benchmark-commands) +5. [Lessons](#5-lessons) + +--- + +## 1. Native SM120 FP4 x FP4 GEMM + +The default FP4 operator is weight-only: `A` remains FP16/BF16, `B` is +dequantized to the activation type, and cuBLAS computes the GEMM. On Blackwell +SM120, CUTLASS also supports native block-scaled NVFP4 x NVFP4 GEMM. The native +path was added behind: + +```bash +ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 +``` + +Native path conditions: + +- SM120 device, +- `M > 8` because decode GEMV has priority, +- `block_size == 16`, +- `K % 32 == 0`, +- `N % 32 == 0`. + +The path quantizes `A` to packed NVFP4 E2M1, creates per-16-block E4M3 +activation scales, uses SM120-swizzled B scales, and launches CUTLASS +block-scaled GEMM. + +Accuracy is compared against an activation-quantized FP4 reference in the Python +harness. This is intentional: the native path is not bitwise-equivalent to the +default weight-only FP4 path because it also quantizes activations. + +Representative result on Blackwell GPU 0 for `M=16,N=11008,K=4096,fp16`: + +| Path | Mean latency | TFLOP/s | Accuracy reference | +|------|--------------|---------|--------------------| +| Default dequant + cuBLAS | about 0.71 ms | about 2.0 | weight-only FP4 reference | +| Native SM120 FP4 x FP4, with prepacked B scales | about 0.15-0.16 ms | about 9.1-9.5 | activation-quantized FP4 reference | + +--- + +## 2. Prepacking SM120 Swizzled B Scales + +The native CUTLASS GEMM expects B scales in an SM120 swizzled scale layout, not +the operator's original `[N, K / 16]` row-major scale tensor. Repacking the scale +tensor on every run was measurable overhead, so `PrePack` now repacks initializer +`weight_scale` once into `b_scale_prepacked_`. + +Important detail: `PrePack` does **not** mark `weight_scale` as removable. The +original unswizzled scale input is still needed by: + +- decode GEMV, +- default dequantize + cuBLAS fallback, +- dynamic cases where native SM120 is not selected. + +Measured effect on the same representative prefill shape: + +| Variant | Mean latency | +|---------|--------------| +| Native SM120 with per-run B-scale repack | about 0.19-0.20 ms | +| Native SM120 with `PrePack` cached B-scale repack | about 0.15-0.16 ms | + +This optimization is kept. + +--- + +## 3. Rejected Experiment - Reuse Swizzled B Scales for Decode GEMV + +Question tested: can the decode GEMV use the same prepacked SM120 swizzled B +scale buffer and avoid reading the original unswizzled `weight_scale`? + +Implementation attempted: + +- Add a sibling GEMV launcher that accepts `b_scale_prepacked_`. +- Add a swizzled-scale accessor equivalent to the SM120 CUTLASS scale layout. +- Route decode GEMV to this launcher when `b_scale_prepacked_` exists. + +Result: correct, but slower. + +Reason: the existing decode GEMV maps one warp to one output column. For that +access pattern, the original layout gives contiguous per-column scale loads: + +``` +weight_scale[col * k_blocks + kb] +``` + +The SM120 swizzled layout is optimized for tiled block-scaled GEMM. For one +fixed output column, consecutive K-block scale loads jump through memory. Even +after precomputing the row base and reducing accessor arithmetic, the swizzled +scale layout was slower for decode. + +Representative `M=1,N=11008,K=4096,fp16` result: + +| Decode variant | Mean latency | Notes | +|----------------|--------------|-------| +| Original GEMV, unswizzled scales | about 0.116 ms | contiguous scale row | +| Swizzled-scale GEMV attempt | about 0.138 ms | strided scale access | + +Representative `M=8,N=11008,K=4096,fp16` result: + +| Decode variant | Mean latency | Notes | +|----------------|--------------|-------| +| Original GEMV, unswizzled scales | about 0.382 ms | contiguous scale row | +| Swizzled-scale GEMV attempt | about 0.507 ms | strided scale access | + +Decision: remove the swizzled-scale GEMV path and keep decode on the original +unswizzled scale layout. + +--- + +## 4. Benchmark Commands + +Provider rebuild and Python-provider sync: + +```bash +cmake --build build/cu130/Release --target onnxruntime_providers_cuda --parallel +cp build/cu130/Release/libonnxruntime_providers_cuda.so \ + build/cu130/Release/onnxruntime/capi/libonnxruntime_providers_cuda.so +cp build/cu130/Release/libonnxruntime_providers_cuda.so \ + build/cu130/Release/build/lib/onnxruntime/capi/libonnxruntime_providers_cuda.so +``` + +Decode benchmarks: + +```bash +cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + --op fp4 --activation-dtype fp16 --m 1 --n 11008 --k 4096 --warmup 100 --repeat 500 + +cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + --op fp4 --activation-dtype fp16 --m 8 --n 11008 --k 4096 --warmup 100 --repeat 500 +``` + +Native prefill benchmark: + +```bash +cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + --op fp4 --activation-dtype fp16 --m 16 --n 11008 --k 4096 --warmup 50 --repeat 200 +``` + +Focused C++ tests: + +```bash +CUDA_VISIBLE_DEVICES=0 build/cu130/Release/onnxruntime_provider_test \ + --gtest_filter='MatMulBlockScaledFp4OpTest.*' +``` + +--- + +## 5. Lessons + +- Keep the native SM120 swizzled scale layout for native GEMM only. +- Keep decode GEMV on the original `[N, K / 16]` scale layout. +- Prepacking can still cache the native GEMM swizzled scale buffer, but it must + leave the original `weight_scale` input available. +- When a native path changes arithmetic semantics by quantizing activations, + validate against an activation-quantized reference, not the weight-only FP4 + reference. diff --git a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc index da7ef35d25052..58807c7197844 100644 --- a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc @@ -180,6 +180,7 @@ class CUDA_ONNX_OP_TYPED_CLASS_NAME(1, float_float_MLFloat16, SimplifiedLayerNor class CUDA_ONNX_OP_TYPED_CLASS_NAME(1, MLFloat16_float_float, SimplifiedLayerNormalization); class CUDA_ONNX_OP_TYPED_CLASS_NAME(1, BFloat16_float_BFloat16, SimplifiedLayerNormalization); class CUDA_MS_OP_CLASS_NAME(1, Inverse); +class CUDA_MS_OP_CLASS_NAME(1, MatMulBlockScaledFp4); class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, MatMulNBits); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, MatMulNBits); class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, MatMulNBits); @@ -441,6 +442,7 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc new file mode 100644 index 0000000000000..45f35788cea24 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -0,0 +1,284 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/cuda/math/matmul_block_scaled_fp4.h" + +#include +#include + +#include "core/common/safeint.h" +#include "core/providers/cuda/cuda_common.h" +#include "core/providers/cuda/shared_inc/fpgeneric.h" +#include "core/providers/cpu/math/matmul_helper.h" +#include "core/platform/env_var_utils.h" + +namespace onnxruntime::contrib::cuda { +using namespace onnxruntime::cuda; + +ONNX_OPERATOR_KERNEL_EX( + MatMulBlockScaledFp4, + kMSDomain, + 1, + kCudaExecutionProvider, + (*KernelDefBuilder::Create()) + .TypeConstraint("T", BuildKernelDefConstraints()) + .TypeConstraint("T1", BuildKernelDefConstraints()) + .TypeConstraint("T2", BuildKernelDefConstraints()) + .TypeConstraint("T3", BuildKernelDefConstraints()), + MatMulBlockScaledFp4); + +namespace { + +constexpr int kWeightScaleInputIndex = 2; + +bool IsNativeSm120Fp4Enabled() { + return ParseEnvironmentVariableWithDefault("ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120", false); +} + +int64_t RoundUp(int64_t value, int64_t alignment) { + return ((value + alignment - 1) / alignment) * alignment; +} + +} // namespace + +MatMulBlockScaledFp4::MatMulBlockScaledFp4(const OpKernelInfo& info) : CudaKernel(info) { + ORT_ENFORCE(info.GetAttr("K", &K_).IsOK()); + ORT_ENFORCE(info.GetAttr("N", &N_).IsOK()); + block_size_ = info.GetAttrOrDefault("block_size", static_cast(16)); + ORT_ENFORCE(K_ > 0, "K must be positive, got ", K_); + ORT_ENFORCE(N_ > 0, "N must be positive, got ", N_); + ORT_ENFORCE(block_size_ > 0, "block_size must be positive, got ", block_size_); + ORT_ENFORCE(K_ % 2 == 0, "K must be even for packed NVFP4 weights, got ", K_); + sm_ = GetDeviceProp().major * 10 + GetDeviceProp().minor; +} + +Status MatMulBlockScaledFp4::PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, + bool& is_packed, PrePackedWeights* /*prepacked_weights*/) { + is_packed = false; + +#if defined(ORT_ENABLE_BLOCKQUANT_SM120) + if (input_idx != kWeightScaleInputIndex || !IsNativeSm120Fp4Enabled() || sm_ < 120 || sm_ >= 130 || + block_size_ != 16 || K_ % 32 != 0 || N_ % 32 != 0) { + return Status::OK(); + } + + const int64_t k_blocks = K_ / 16; + ORT_RETURN_IF_NOT(tensor.Shape().Size() >= N_ * k_blocks, + "weight_scale tensor is too small; expected at least ", N_ * k_blocks, " E4M3 scales."); + + const int64_t rounded_k_blocks = RoundUp(k_blocks, 4); + const int64_t rounded_n = RoundUp(N_, 128); + b_scale_prepacked_ = IAllocator::MakeUniquePtr( + alloc, SafeInt(rounded_n) * SafeInt(rounded_k_blocks), true); + + cudaStream_t stream = cudaStreamLegacy; + const void* weight_scale = tensor.DataRaw(); + IAllocatorUniquePtr weight_scale_device; + if (tensor.Location().device.Type() != OrtDevice::GPU) { + const size_t weight_scale_bytes = SafeInt(N_) * SafeInt(k_blocks); + weight_scale_device = IAllocator::MakeUniquePtr(alloc, weight_scale_bytes, true); + CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(weight_scale_device.get(), weight_scale, weight_scale_bytes, + cudaMemcpyDefault, stream)); + weight_scale = weight_scale_device.get(); + } + + ORT_RETURN_IF_ERROR(LaunchRepackWeightScaleNvFp4ForNativeSm120( + b_scale_prepacked_.get(), weight_scale, SafeInt(N_), SafeInt(K_), SafeInt(block_size_), stream)); + CUDA_RETURN_IF_ERROR(cudaStreamSynchronize(stream)); +#else + ORT_UNUSED_PARAMETER(tensor); + ORT_UNUSED_PARAMETER(input_idx); + ORT_UNUSED_PARAMETER(alloc); +#endif + + return Status::OK(); +} + +template +Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { + typedef typename ToCudaType::MappedType CudaT; + + const Tensor* a = context->Input(0); + const Tensor* b = context->Input(1); + const Tensor* weight_scale = context->Input(2); + const Tensor* weight_scale_2 = context->Input(3); + const Tensor* input_scale = context->Input(4); // optional + const Tensor* bias = context->Input(5); // optional + + const auto& a_shape = a->Shape(); + ORT_ENFORCE(a_shape.NumDimensions() >= 1, "A must have rank at least 1."); + ORT_ENFORCE(a_shape[a_shape.NumDimensions() - 1] == K_, + "A's last dimension (", a_shape[a_shape.NumDimensions() - 1], ") must equal K (", K_, ")."); + + const int64_t k_packed = K_ / 2; + const int64_t k_blocks = (K_ + block_size_ - 1) / block_size_; + ORT_ENFORCE(b->Shape().Size() >= N_ * k_packed, + "B tensor is too small; expected at least ", N_ * k_packed, " packed bytes."); + ORT_ENFORCE(weight_scale->Shape().Size() >= N_ * k_blocks, + "weight_scale tensor is too small; expected at least ", N_ * k_blocks, " E4M3 scales."); + ORT_ENFORCE(weight_scale_2->Shape().Size() == 1, "weight_scale_2 must be a scalar."); + if (input_scale != nullptr) { + ORT_ENFORCE(input_scale->Shape().Size() == 1, "input_scale must be a scalar."); + // input_scale is used only by the opt-in native NVFP4 x NVFP4 path. The default + // weight-only FP16/BF16 activation path keeps full-precision activations. + } + if (bias != nullptr) { + ORT_ENFORCE(bias->Shape().Size() == N_, "bias must have shape [N]."); + } + + constexpr bool transa = false; + constexpr bool transb = true; + MatMulComputeHelper helper; + TensorShape b_logical_shape({N_, K_}); + ORT_RETURN_IF_ERROR(helper.Compute(a_shape, b_logical_shape, transa, transb)); + + Tensor* Y = context->Output(0, helper.OutputShape()); + if (Y->Shape().Size() == 0) { + return Status::OK(); + } + + const int m_i = SafeInt(helper.M()); + const int n_i = SafeInt(helper.N()); + const int k_i = SafeInt(helper.K()); + + // Decode fast path: for small M (autoregressive generation) this is a memory-bound GEMV. + // A fused warp-per-column kernel reads the packed NVFP4 weight directly, avoiding both the + // [N, K] dequant scratch buffer and the cuBLAS GEMM (which is underutilized at M == 1). + constexpr int kGemvMaxM = 8; + if (m_i > 0 && m_i <= kGemvMaxM && block_size_ == 16 && (k_i % 32 == 0)) { + return LaunchMatMulBlockScaledFp4Gemv( + Y->MutableDataRaw(), + a->DataRaw(), + b->DataRaw(), + weight_scale->DataRaw(), + weight_scale_2->Data(), + bias != nullptr ? bias->DataRaw() : nullptr, + m_i, + n_i, + k_i, + SafeInt(block_size_), + std::is_same::value, + Stream(context)); + } + +#if defined(ORT_ENABLE_BLOCKQUANT_SM120) + if (IsNativeSm120Fp4Enabled() && sm_ >= 120 && sm_ < 130 && block_size_ == 16 && + (k_i % 32 == 0) && (n_i % 32 == 0)) { + constexpr int64_t kScaleVectorSize = 16; + const int64_t k_scale_blocks = RoundUp(k_i / kScaleVectorSize, 4); + const int64_t rounded_m = RoundUp(m_i, 128); + const int64_t rounded_n = RoundUp(n_i, 128); + + auto a_packed = GetScratchBuffer(SafeInt(m_i) * SafeInt(k_i / 2), + context->GetComputeStream()); + auto a_scale = GetScratchBuffer(SafeInt(rounded_m) * SafeInt(k_scale_blocks), + context->GetComputeStream()); + IAllocatorUniquePtr b_scale; + const void* b_scale_data = b_scale_prepacked_.get(); + if (b_scale_data == nullptr) { + b_scale = GetScratchBuffer(SafeInt(rounded_n) * SafeInt(k_scale_blocks), + context->GetComputeStream()); + ORT_RETURN_IF_ERROR(LaunchRepackWeightScaleNvFp4ForNativeSm120( + b_scale.get(), weight_scale->DataRaw(), n_i, k_i, SafeInt(block_size_), Stream(context))); + b_scale_data = b_scale.get(); + } + auto alpha = GetScratchBuffer(1, context->GetComputeStream()); + const size_t workspace_size = GetMatMulBlockScaledFp4NativeSm120WorkspaceSize( + m_i, n_i, k_i, std::is_same::value); + auto workspace = GetScratchBuffer(workspace_size, context->GetComputeStream()); + + ORT_RETURN_IF_ERROR(LaunchMatMulBlockScaledFp4NativeSm120( + Y->MutableDataRaw(), + a->DataRaw(), + b->DataRaw(), + weight_scale->DataRaw(), + weight_scale_2->Data(), + input_scale != nullptr ? input_scale->Data() : nullptr, + a_packed.get(), + a_scale.get(), + b_scale_data, + alpha.get(), + m_i, + n_i, + k_i, + SafeInt(block_size_), + std::is_same::value, + workspace.get(), + workspace_size, + Stream(context))); + + if (bias != nullptr) { + ORT_RETURN_IF_ERROR(LaunchAddBiasNvFp4( + Y->MutableDataRaw(), + bias->DataRaw(), + m_i, + n_i, + std::is_same::value, + Stream(context))); + } + + return Status::OK(); + } +#endif + + // Dequantize the packed NVFP4 weight into a scratch [N, K] buffer of the activation type. + IAllocatorUniquePtr b_dequant = GetScratchBuffer(SafeInt(N_) * SafeInt(K_), + context->GetComputeStream()); + ORT_RETURN_IF_ERROR(LaunchDequantizeNvFp4( + b_dequant.get(), + b->DataRaw(), + weight_scale->DataRaw(), + weight_scale_2->Data(), + SafeInt(N_), + SafeInt(K_), + SafeInt(block_size_), + std::is_same::value, + Stream(context))); + + const CudaT alpha = ToCudaType::FromFloat(1.f); + const CudaT zero = ToCudaType::FromFloat(0.f); + + CUBLAS_RETURN_IF_ERROR(cublasGemmHelper( + GetCublasHandle(context), + CUBLAS_OP_T, // transB: dequantized weight is [N, K] row-major == K-major [K, N] + CUBLAS_OP_N, // transA + n_i, + m_i, + k_i, + &alpha, + b_dequant.get(), + helper.Ldb(transb), + reinterpret_cast(a->DataRaw()), + helper.Lda(transa), + &zero, + reinterpret_cast(Y->MutableDataRaw()), + helper.Ldc(), + GetDeviceProp(), + UseTF32())); + + if (bias != nullptr) { + ORT_RETURN_IF_ERROR(LaunchAddBiasNvFp4( + Y->MutableDataRaw(), + bias->DataRaw(), + m_i, + n_i, + std::is_same::value, + Stream(context))); + } + + return Status::OK(); +} + +Status MatMulBlockScaledFp4::ComputeInternal(OpKernelContext* context) const { + const Tensor* a = context->Input(0); + if (a->IsDataType()) { + return ComputeImpl(context); + } + if (a->IsDataType()) { + return ComputeImpl(context); + } + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, + "MatMulBlockScaledFp4 only supports FP16 or BF16 activations."); +} + +} // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu new file mode 100644 index 0000000000000..26e694e6ce0a0 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu @@ -0,0 +1,349 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/cuda/math/matmul_block_scaled_fp4.h" + +#include +#include +#include + +#if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 +#include +#include +#endif + +#include "core/providers/cuda/cuda_common.h" + +namespace onnxruntime::contrib::cuda { + +#if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 + +namespace { + +template +__device__ __forceinline__ T FromFloat(float v); + +template <> +__device__ __forceinline__ half FromFloat(float v) { + return __float2half(v); +} + +template <> +__device__ __forceinline__ nv_bfloat16 FromFloat(float v) { + return __float2bfloat16(v); +} + +template +__device__ __forceinline__ float ToFloat(T v); + +template <> +__device__ __forceinline__ float ToFloat(half v) { + return __half2float(v); +} + +template <> +__device__ __forceinline__ float ToFloat(nv_bfloat16 v) { + return __bfloat162float(v); +} + +template +__global__ void DequantizeNvFp4Kernel(T* __restrict__ out, + const uint8_t* __restrict__ b_packed, + const uint8_t* __restrict__ weight_scale, + const float* __restrict__ weight_scale_2, + int n, + int k, + int k_blocks, + int block_size) { + const int half_k = k >> 1; + const long long total = static_cast(n) * half_k; + const long long idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (idx >= total) { + return; + } + + const int row = static_cast(idx / half_k); + const int pair = static_cast(idx - static_cast(row) * half_k); + const int k0 = pair << 1; + + const uint8_t packed = b_packed[idx]; + const __half2_raw hr = __nv_cvt_fp4x2_to_halfraw2(static_cast<__nv_fp4x2_storage_t>(packed), __NV_E2M1); + const __half2 h2 = __half2(hr); + const float2 v = __half22float2(h2); + + const float g = *weight_scale_2; + const int blk0 = k0 / block_size; + const int blk1 = (k0 + 1) / block_size; + const float s0 = __half2float(__nv_cvt_fp8_to_halfraw( + static_cast<__nv_fp8_storage_t>(weight_scale[row * k_blocks + blk0]), __NV_E4M3)) * + g; + const float s1 = __half2float(__nv_cvt_fp8_to_halfraw( + static_cast<__nv_fp8_storage_t>(weight_scale[row * k_blocks + blk1]), __NV_E4M3)) * + g; + + const long long out_base = static_cast(row) * k + k0; + out[out_base] = FromFloat(v.x * s0); + out[out_base + 1] = FromFloat(v.y * s1); +} + +template +__global__ void AddBiasKernel(T* __restrict__ y, const T* __restrict__ bias, int m, int n) { + const long long total = static_cast(m) * n; + const long long idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (idx >= total) { + return; + } + const int col = static_cast(idx % n); + y[idx] = FromFloat(ToFloat(y[idx]) + ToFloat(bias[col])); +} + +// ----------------------------------------------------------------------------- +// Fused NVFP4 weight-only GEMV fast path for the decode phase (small M). +// +// Each warp computes one output element Y[row, col]. The 32 lanes cooperatively +// reduce over K reading the packed NVFP4 weight directly (two E2M1 values per +// byte) with 16-byte coalesced loads, so the weight is streamed exactly once and +// no [N, K] dequantized buffer is materialized. Each lane consumes 32 contiguous +// K elements = 16 packed bytes, which span exactly two 16-element blocks; the two +// per-block E4M3 scales are folded in per half. The global fp32 scale is applied +// once after the warp reduction. Runs on any architecture with NVFP4 conversion +// intrinsics (CUDA >= 12.8), including SM90 and SM120. +template +__device__ __forceinline__ void LoadFp4Gemv32A(const T* ptr, float (&out)[32]); + +template <> +__device__ __forceinline__ void LoadFp4Gemv32A(const half* ptr, float (&out)[32]) { + const uint4* p = reinterpret_cast(ptr); +#pragma unroll + for (int j = 0; j < 4; ++j) { + const uint4 raw = p[j]; + const half* v = reinterpret_cast(&raw); +#pragma unroll + for (int i = 0; i < 8; ++i) { + out[j * 8 + i] = __half2float(v[i]); + } + } +} + +template <> +__device__ __forceinline__ void LoadFp4Gemv32A(const nv_bfloat16* ptr, float (&out)[32]) { + const uint4* p = reinterpret_cast(ptr); +#pragma unroll + for (int j = 0; j < 4; ++j) { + const uint4 raw = p[j]; + const nv_bfloat16* v = reinterpret_cast(&raw); +#pragma unroll + for (int i = 0; i < 8; ++i) { + out[j * 8 + i] = __bfloat162float(v[i]); + } + } +} + +template +__global__ void MatMulBlockScaledFp4GemvKernel(T* __restrict__ y, + const T* __restrict__ a, + const uint8_t* __restrict__ b_packed, + const uint8_t* __restrict__ weight_scale, + const float* __restrict__ weight_scale_2, + const T* __restrict__ bias, + int m, + int n, + int k, + int k_blocks) { + const int lane = threadIdx.x; // 0..31 + const int col = blockIdx.x * blockDim.y + threadIdx.y; // n + const int row = blockIdx.y; // m + if (row >= m || col >= n) { + return; + } + + const T* a_row = a + static_cast(row) * k; + const uint8_t* b_row = b_packed + static_cast(col) * (k >> 1); + const uint8_t* ws_row = weight_scale + static_cast(col) * k_blocks; + + constexpr int kBlockSize = 16; + constexpr int kElemsPerLane = 32; // two 16-element blocks + const int stride = 32 * kElemsPerLane; // 1024 elements per warp iteration + + float acc = 0.0f; + for (int base = 0; base < k; base += stride) { + const int koff = base + lane * kElemsPerLane; + if (koff < k) { + const uint4 packed = *reinterpret_cast(b_row + (koff >> 1)); + const uint8_t* bytes = reinterpret_cast(&packed); + float b_vals[32]; +#pragma unroll + for (int i = 0; i < 16; ++i) { + const __half2_raw hr = __nv_cvt_fp4x2_to_halfraw2( + static_cast<__nv_fp4x2_storage_t>(bytes[i]), __NV_E2M1); + const float2 f = __half22float2(__half2(hr)); + b_vals[i * 2] = f.x; + b_vals[i * 2 + 1] = f.y; + } + + float a_vals[32]; + LoadFp4Gemv32A(a_row + koff, a_vals); + + const int kb0 = koff / kBlockSize; + const int kb1 = kb0 + 1; + float p0 = 0.0f; + float p1 = 0.0f; +#pragma unroll + for (int i = 0; i < 16; ++i) { + p0 += a_vals[i] * b_vals[i]; + } +#pragma unroll + for (int i = 16; i < 32; ++i) { + p1 += a_vals[i] * b_vals[i]; + } + const float s0 = __half2float(__nv_cvt_fp8_to_halfraw( + static_cast<__nv_fp8_storage_t>(ws_row[kb0]), __NV_E4M3)); + const float s1 = __half2float(__nv_cvt_fp8_to_halfraw( + static_cast<__nv_fp8_storage_t>(ws_row[kb1]), __NV_E4M3)); + acc += p0 * s0 + p1 * s1; + } + } + +#pragma unroll + for (int offset = 16; offset > 0; offset >>= 1) { + acc += __shfl_down_sync(0xffffffffu, acc, offset); + } + if (lane == 0) { + float result = acc * (*weight_scale_2); + if (bias != nullptr) { + result += ToFloat(bias[col]); + } + y[static_cast(row) * n + col] = FromFloat(result); + } +} + +} // namespace + +#endif // CUDA_VERSION >= 12080 + +Status LaunchDequantizeNvFp4(void* b_dequant, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + int n, + int k, + int block_size, + bool is_bf16, + cudaStream_t stream) { +#if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 + const int half_k = k >> 1; + const long long total = static_cast(n) * half_k; + if (total == 0) { + return Status::OK(); + } + const int k_blocks = (k + block_size - 1) / block_size; + constexpr int kThreads = 256; + const int blocks = static_cast((total + kThreads - 1) / kThreads); + const uint8_t* bp = reinterpret_cast(b_packed); + const uint8_t* ws = reinterpret_cast(weight_scale); + + if (is_bf16) { + DequantizeNvFp4Kernel<<>>( + reinterpret_cast(b_dequant), bp, ws, weight_scale_2, n, k, k_blocks, block_size); + } else { + DequantizeNvFp4Kernel<<>>( + reinterpret_cast(b_dequant), bp, ws, weight_scale_2, n, k, k_blocks, block_size); + } + return CUDA_CALL(cudaGetLastError()); +#else + ORT_UNUSED_PARAMETER(b_dequant); + ORT_UNUSED_PARAMETER(b_packed); + ORT_UNUSED_PARAMETER(weight_scale); + ORT_UNUSED_PARAMETER(weight_scale_2); + ORT_UNUSED_PARAMETER(n); + ORT_UNUSED_PARAMETER(k); + ORT_UNUSED_PARAMETER(block_size); + ORT_UNUSED_PARAMETER(is_bf16); + ORT_UNUSED_PARAMETER(stream); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockScaledFp4 requires CUDA 12.8 or newer for NVFP4 support."); +#endif +} + +Status LaunchAddBiasNvFp4(void* y, + const void* bias, + int m, + int n, + bool is_bf16, + cudaStream_t stream) { +#if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 + const long long total = static_cast(m) * n; + if (total == 0) { + return Status::OK(); + } + constexpr int kThreads = 256; + const int blocks = static_cast((total + kThreads - 1) / kThreads); + if (is_bf16) { + AddBiasKernel<<>>( + reinterpret_cast(y), reinterpret_cast(bias), m, n); + } else { + AddBiasKernel<<>>( + reinterpret_cast(y), reinterpret_cast(bias), m, n); + } + return CUDA_CALL(cudaGetLastError()); +#else + ORT_UNUSED_PARAMETER(y); + ORT_UNUSED_PARAMETER(bias); + ORT_UNUSED_PARAMETER(m); + ORT_UNUSED_PARAMETER(n); + ORT_UNUSED_PARAMETER(is_bf16); + ORT_UNUSED_PARAMETER(stream); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockScaledFp4 requires CUDA 12.8 or newer for NVFP4 support."); +#endif +} + +Status LaunchMatMulBlockScaledFp4Gemv(void* y, + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const void* bias, + int m, + int n, + int k, + int block_size, + bool is_bf16, + cudaStream_t stream) { +#if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 + if (m <= 0 || n <= 0 || k <= 0) { + return Status::OK(); + } + const int k_blocks = (k + block_size - 1) / block_size; + constexpr int kWarpsPerBlock = 8; + const dim3 threads{32, kWarpsPerBlock}; + const dim3 blocks{static_cast((n + kWarpsPerBlock - 1) / kWarpsPerBlock), + static_cast(m)}; + const uint8_t* bp = reinterpret_cast(b_packed); + const uint8_t* ws = reinterpret_cast(weight_scale); + if (is_bf16) { + MatMulBlockScaledFp4GemvKernel<<>>( + reinterpret_cast(y), reinterpret_cast(a), bp, ws, weight_scale_2, + reinterpret_cast(bias), m, n, k, k_blocks); + } else { + MatMulBlockScaledFp4GemvKernel<<>>( + reinterpret_cast(y), reinterpret_cast(a), bp, ws, weight_scale_2, + reinterpret_cast(bias), m, n, k, k_blocks); + } + return CUDA_CALL(cudaGetLastError()); +#else + ORT_UNUSED_PARAMETER(y); + ORT_UNUSED_PARAMETER(a); + ORT_UNUSED_PARAMETER(b_packed); + ORT_UNUSED_PARAMETER(weight_scale); + ORT_UNUSED_PARAMETER(weight_scale_2); + ORT_UNUSED_PARAMETER(bias); + ORT_UNUSED_PARAMETER(m); + ORT_UNUSED_PARAMETER(n); + ORT_UNUSED_PARAMETER(k); + ORT_UNUSED_PARAMETER(block_size); + ORT_UNUSED_PARAMETER(is_bf16); + ORT_UNUSED_PARAMETER(stream); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockScaledFp4 requires CUDA 12.8 or newer for NVFP4 support."); +#endif +} + +} // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h new file mode 100644 index 0000000000000..3da8d52ac8e6c --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h @@ -0,0 +1,109 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/providers/cuda/cuda_kernel.h" + +namespace onnxruntime::contrib::cuda { + +// Weight-only NVFP4 (E2M1) matrix multiplication. +// +// The weight tensor B is stored as packed NVFP4: two E2M1 values per byte (low nibble first), +// with a per-16-block E4M3 scale (weight_scale) and a single global fp32 scale (weight_scale_2). +// The weight is dequantized to the activation type (FP16/BF16) and multiplied with the FP16/BF16 +// activation via cuBLAS. This path works on any CUDA architecture (including Hopper/SM90) because +// it does not rely on native NVFP4 block-scaled tensor cores (SM100/SM120 only). +class MatMulBlockScaledFp4 final : public onnxruntime::cuda::CudaKernel { + public: + explicit MatMulBlockScaledFp4(const OpKernelInfo& info); + + Status PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, + bool& is_packed, PrePackedWeights* prepacked_weights) override; + + Status ComputeInternal(OpKernelContext* context) const override; + + private: + template + Status ComputeImpl(OpKernelContext* context) const; + + int64_t K_; + int64_t N_; + int64_t block_size_; + int sm_{0}; + IAllocatorUniquePtr b_scale_prepacked_; +}; + +// Dequantizes NVFP4 (E2M1) weights with per-block E4M3 scales and a global fp32 scale into +// FP16/BF16. b_packed is [N, K/2] uint8 (two E2M1 values per byte, low nibble first), +// weight_scale is [N, ceil(K/block_size)] uint8 (raw E4M3 bytes), weight_scale_2 is a device +// fp32 scalar. Output b_dequant is [N, K] in the activation type (is_bf16 selects BF16 vs FP16). +Status LaunchDequantizeNvFp4(void* b_dequant, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + int n, + int k, + int block_size, + bool is_bf16, + cudaStream_t stream); + +// Adds a per-column bias of shape [N] to a [M, N] row-major output in place. +Status LaunchAddBiasNvFp4(void* y, + const void* bias, + int m, + int n, + bool is_bf16, + cudaStream_t stream); + +// Fused NVFP4 weight-only GEMV fast path for the decode phase (small M). Reads the packed +// NVFP4 weight directly (no [N, K] dequant buffer). a is [M, K] activation (FP16/BF16), +// b_packed is [N, K/2] uint8 (two E2M1 values per byte), weight_scale is [N, ceil(K/block_size)] +// uint8 (raw E4M3 bytes), weight_scale_2 is a device fp32 scalar, bias is an optional [N] vector +// (may be null). Output y is [M, N] in the activation type. Requires block_size == 16 and +// k % 32 == 0. Runs on any architecture with NVFP4 conversion intrinsics (CUDA >= 12.8). +Status LaunchMatMulBlockScaledFp4Gemv(void* y, + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const void* bias, + int m, + int n, + int k, + int block_size, + bool is_bf16, + cudaStream_t stream); + +// Native Blackwell SM120 NVFP4 x NVFP4 GEMM path. The caller provides scratch buffers for +// packed activation FP4, swizzled A/B scale tensors, alpha, and CUTLASS workspace. A is [M, K] +// FP16/BF16, B is [N, K/2] packed NVFP4, weight_scale is [N, K/16] E4M3, and Y is [M, N] +// FP16/BF16. Requires block_size == 16, K % 32 == 0, and N % 32 == 0. +Status LaunchRepackWeightScaleNvFp4ForNativeSm120(void* b_scale, + const void* weight_scale, + int n, + int k, + int block_size, + cudaStream_t stream); + +Status LaunchMatMulBlockScaledFp4NativeSm120(void* y, + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const float* input_scale, + void* a_packed, + void* a_scale, + const void* b_scale, + float* alpha, + int m, + int n, + int k, + int block_size, + bool is_bf16, + void* workspace, + size_t workspace_size, + cudaStream_t stream); +size_t GetMatMulBlockScaledFp4NativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16); + +} // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu new file mode 100644 index 0000000000000..9ed24f97a39e2 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu @@ -0,0 +1,360 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/cuda/math/matmul_block_scaled_fp4.h" + +#include "core/providers/cuda/cuda_common.h" + +#include +#include +#include +#include + +#include "cutlass/bfloat16.h" +#include "cutlass/cutlass.h" +#include "cutlass/epilogue/collective/collective_builder.hpp" +#include "cutlass/gemm/collective/collective_builder.hpp" +#include "cutlass/gemm/device/gemm_universal_adapter.h" +#include "cutlass/gemm/kernel/gemm_universal.hpp" +#include "cutlass/half.h" +#include "cutlass/numeric_types.h" +#include "cutlass/util/packed_stride.hpp" + +namespace onnxruntime::contrib::cuda { +namespace { + +using namespace cute; + +constexpr int kScaleVectorSize = 16; +constexpr int kAlignment = 32; + +template +__device__ __forceinline__ float LoadAsFloat(const T* data, int index); + +template <> +__device__ __forceinline__ float LoadAsFloat(const half* data, int index) { + return __half2float(data[index]); +} + +template <> +__device__ __forceinline__ float LoadAsFloat(const nv_bfloat16* data, int index) { + return __bfloat162float(data[index]); +} + +__device__ __forceinline__ int SwizzledScaleOffset(int row, int k_block, int num_k_tiles) { + const int row_tile = row >> 7; + const int outer_row = row & 31; + const int inner_row = (row >> 5) & 3; + const int k_tile = k_block >> 2; + const int inner_k = k_block & 3; + return ((row_tile * num_k_tiles + k_tile) << 9) | (outer_row << 4) | (inner_row << 2) | inner_k; +} + +__device__ __forceinline__ uint8_t FloatToE4m3(float value) { + __nv_fp8_e4m3 converted(value); + uint8_t raw; + reinterpret_cast<__nv_fp8_e4m3&>(raw) = converted; + return raw; +} + +__device__ __forceinline__ float E4m3ToFloat(uint8_t raw) { + return __half2float(__nv_cvt_fp8_to_halfraw(static_cast<__nv_fp8_storage_t>(raw), __NV_E4M3)); +} + +template +__global__ void QuantizeActivationNvFp4Kernel(const T* __restrict__ a, + const float* __restrict__ input_scale, + uint8_t* __restrict__ a_packed, + uint8_t* __restrict__ a_scale, + float* __restrict__ alpha, + const float* __restrict__ weight_scale_2, + int m, + int k, + int rounded_k_blocks) { + if (blockIdx.x == 0 && threadIdx.x == 0) { + const float activation_global_scale = input_scale != nullptr ? input_scale[0] : 1.0f; + alpha[0] = weight_scale_2[0] / activation_global_scale; + } + + const int row = static_cast(blockIdx.y); + const int k_block = static_cast(blockIdx.x); + if (row >= m || k_block >= k / kScaleVectorSize) { + return; + } + + const int k_base = k_block * kScaleVectorSize; + float values[kScaleVectorSize]; + float max_abs = 0.0f; +#pragma unroll + for (int offset = 0; offset < kScaleVectorSize; ++offset) { + const float value = LoadAsFloat(a, row * k + k_base + offset); + values[offset] = value; + max_abs = fmaxf(max_abs, fabsf(value)); + } + + const float activation_global_scale = input_scale != nullptr ? input_scale[0] : 1.0f; + const uint8_t raw_scale = FloatToE4m3(fmaxf(max_abs / 6.0f, 1.0f / 1024.0f) * activation_global_scale); + a_scale[SwizzledScaleOffset(row, k_block, rounded_k_blocks / 4)] = raw_scale; + const float local_scale = E4m3ToFloat(raw_scale) / activation_global_scale; + +#pragma unroll + for (int pair = 0; pair < kScaleVectorSize / 2; ++pair) { + const float2 scaled = make_float2(values[pair * 2] / local_scale, values[pair * 2 + 1] / local_scale); + a_packed[row * (k / 2) + k_base / 2 + pair] = + static_cast(__nv_cvt_float2_to_fp4x2(scaled, __NV_E2M1, cudaRoundNearest)); + } +} + +__global__ void RepackWeightScaleNvFp4Kernel(const uint8_t* __restrict__ weight_scale, + uint8_t* __restrict__ b_scale, + int n, + int k_blocks, + int rounded_k_blocks) { + const int index = static_cast(blockIdx.x * blockDim.x + threadIdx.x); + const int total = n * rounded_k_blocks; + if (index >= total) { + return; + } + + const int row = index / rounded_k_blocks; + const int k_block = index - row * rounded_k_blocks; + const uint8_t scale = k_block < k_blocks ? weight_scale[row * k_blocks + k_block] : 0; + b_scale[SwizzledScaleOffset(row, k_block, rounded_k_blocks / 4)] = scale; +} + +struct Fp4GemmSm120M256Config { + using KernelSchedule = cutlass::gemm::collective::KernelScheduleAuto; + using EpilogueSchedule = cutlass::epilogue::collective::EpilogueScheduleAuto; + using TileScheduler = void; + using ClusterShape = Shape<_1, _1, _1>; + using MmaTileShape = Shape<_128, _128, _128>; + using PerSmTileShape = Shape<_128, _128, _128>; +}; + +struct Fp4GemmSm120DefaultConfig { + using KernelSchedule = cutlass::gemm::collective::KernelScheduleAuto; + using EpilogueSchedule = cutlass::epilogue::collective::EpilogueScheduleAuto; + using TileScheduler = cutlass::gemm::PersistentScheduler; + using ClusterShape = Shape<_1, _1, _1>; + using MmaTileShape = Shape<_256, _128, _128>; + using PerSmTileShape = Shape<_256, _128, _128>; +}; + +template +struct Fp4GemmSm120 { + using ElementA = cutlass::nv_float4_t; + using LayoutA = cutlass::layout::RowMajor; + using ElementB = cutlass::nv_float4_t; + using LayoutB = cutlass::layout::ColumnMajor; + using ElementC = OutType; + using ElementD = OutType; + using LayoutC = cutlass::layout::RowMajor; + using LayoutD = cutlass::layout::RowMajor; + using ElementAccumulator = float; + using ArchTag = cutlass::arch::Sm120; + using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + ArchTag, OperatorClass, typename Config::PerSmTileShape, typename Config::ClusterShape, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementAccumulator, + ElementC, LayoutC, 128 / cutlass::sizeof_bits::value, + ElementD, LayoutD, 128 / cutlass::sizeof_bits::value, + typename Config::EpilogueSchedule>::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + ArchTag, OperatorClass, + ElementA, LayoutA, kAlignment, + ElementB, LayoutB, kAlignment, + ElementAccumulator, typename Config::MmaTileShape, typename Config::ClusterShape, + cutlass::gemm::collective::StageCountAutoCarveout( + sizeof(typename CollectiveEpilogue::SharedStorage))>, + typename Config::KernelSchedule>::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, CollectiveMainloop, CollectiveEpilogue, typename Config::TileScheduler>; + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; +}; + +template +typename Gemm::Arguments MakeArguments(void* y, + const void* a_packed, + const void* b_packed, + const void* a_scale, + const void* b_scale, + const float* alpha, + int m, + int n, + int k) { + using ElementA = typename Gemm::GemmKernel::ElementA; + using ElementB = typename Gemm::GemmKernel::ElementB; + using ElementD = typename Gemm::GemmKernel::ElementD; + using ElementSF = cutlass::float_ue4m3_t; + using StrideA = typename Gemm::GemmKernel::StrideA; + using StrideB = typename Gemm::GemmKernel::StrideB; + using StrideD = typename Gemm::GemmKernel::StrideD; + using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig; + + constexpr int l = 1; + StrideA stride_a = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(m, k, l)); + StrideB stride_b = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(n, k, l)); + StrideD stride_d = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(m, n, l)); + auto layout_sfa = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n, k, l)); + auto layout_sfb = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, l)); + + typename Gemm::Arguments arguments{ + cutlass::gemm::GemmUniversalMode::kGemm, + {m, n, k, l}, + {reinterpret_cast(a_packed), stride_a, + reinterpret_cast(b_packed), stride_b, + reinterpret_cast(a_scale), layout_sfa, + reinterpret_cast(b_scale), layout_sfb}, + {{}, reinterpret_cast(y), stride_d, reinterpret_cast(y), stride_d}}; + arguments.epilogue.thread.alpha_ptr = alpha; + return arguments; +} + +template +size_t WorkspaceSize(int m, int n, int k) { + auto arguments = MakeArguments(nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, m, n, k); + return Gemm::get_workspace_size(arguments); +} + +template +Status RunGemm(void* y, + const void* a_packed, + const void* b_packed, + const void* a_scale, + const void* b_scale, + const float* alpha, + int m, + int n, + int k, + void* workspace, + cudaStream_t stream) { + auto arguments = MakeArguments(y, a_packed, b_packed, a_scale, b_scale, alpha, m, n, k); + Gemm gemm; + cutlass::Status status = gemm.can_implement(arguments); + ORT_RETURN_IF_NOT(status == cutlass::Status::kSuccess, + "SM120 native FP4 GEMM cannot implement the given problem: ", + cutlassGetStatusString(status)); + status = gemm.initialize(arguments, workspace, stream); + ORT_RETURN_IF_NOT(status == cutlass::Status::kSuccess, + "SM120 native FP4 GEMM initialize failed: ", cutlassGetStatusString(status)); + status = gemm.run(arguments, workspace, stream); + ORT_RETURN_IF_NOT(status == cutlass::Status::kSuccess, + "SM120 native FP4 GEMM run failed: ", cutlassGetStatusString(status)); + return CUDA_CALL(cudaGetLastError()); +} + +bool UseM256Config(int m) { + const auto m_unsigned = static_cast(std::max(m - 1, 1)); + const int next_power_of_two_m = static_cast(1u << (32 - __builtin_clz(m_unsigned))); + return std::max(16, next_power_of_two_m) <= 256; +} + +template +size_t DispatchWorkspaceSize(int m, int n, int k) { + if (UseM256Config(m)) { + return WorkspaceSize::Gemm>(m, n, k); + } + return WorkspaceSize::Gemm>(m, n, k); +} + +template +Status DispatchRunGemm(void* y, + const void* a_packed, + const void* b_packed, + const void* a_scale, + const void* b_scale, + const float* alpha, + int m, + int n, + int k, + void* workspace, + cudaStream_t stream) { + if (UseM256Config(m)) { + return RunGemm::Gemm>( + y, a_packed, b_packed, a_scale, b_scale, alpha, m, n, k, workspace, stream); + } + return RunGemm::Gemm>( + y, a_packed, b_packed, a_scale, b_scale, alpha, m, n, k, workspace, stream); +} + +} // namespace + +size_t GetMatMulBlockScaledFp4NativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16) { + return is_bf16 ? DispatchWorkspaceSize(m, n, k) + : DispatchWorkspaceSize(m, n, k); +} + +Status LaunchRepackWeightScaleNvFp4ForNativeSm120(void* b_scale, + const void* weight_scale, + int n, + int k, + int block_size, + cudaStream_t stream) { + ORT_RETURN_IF_NOT(block_size == kScaleVectorSize, + "SM120 native FP4 GEMM only supports block_size == ", kScaleVectorSize); + ORT_RETURN_IF_NOT(k % kAlignment == 0, "SM120 native FP4 GEMM requires K divisible by ", kAlignment); + ORT_RETURN_IF_NOT(n % kAlignment == 0, "SM120 native FP4 GEMM requires N divisible by ", kAlignment); + + const int k_blocks = k / kScaleVectorSize; + const int rounded_k_blocks = ((k_blocks + 3) / 4) * 4; + constexpr int kThreads = 256; + const int repack_total = n * rounded_k_blocks; + const int repack_blocks = (repack_total + kThreads - 1) / kThreads; + RepackWeightScaleNvFp4Kernel<<>>( + reinterpret_cast(weight_scale), reinterpret_cast(b_scale), n, k_blocks, + rounded_k_blocks); + return CUDA_CALL(cudaGetLastError()); +} + +Status LaunchMatMulBlockScaledFp4NativeSm120(void* y, + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const float* input_scale, + void* a_packed, + void* a_scale, + const void* b_scale, + float* alpha, + int m, + int n, + int k, + int block_size, + bool is_bf16, + void* workspace, + size_t workspace_size, + cudaStream_t stream) { + ORT_UNUSED_PARAMETER(workspace_size); + ORT_UNUSED_PARAMETER(weight_scale); + ORT_RETURN_IF_NOT(block_size == kScaleVectorSize, + "SM120 native FP4 GEMM only supports block_size == ", kScaleVectorSize); + ORT_RETURN_IF_NOT(k % kAlignment == 0, "SM120 native FP4 GEMM requires K divisible by ", kAlignment); + ORT_RETURN_IF_NOT(n % kAlignment == 0, "SM120 native FP4 GEMM requires N divisible by ", kAlignment); + + const int k_blocks = k / kScaleVectorSize; + const int rounded_k_blocks = ((k_blocks + 3) / 4) * 4; + const dim3 quant_grid{static_cast(k_blocks), static_cast(m)}; + if (is_bf16) { + QuantizeActivationNvFp4Kernel<<>>( + reinterpret_cast(a), input_scale, reinterpret_cast(a_packed), + reinterpret_cast(a_scale), alpha, weight_scale_2, m, k, rounded_k_blocks); + } else { + QuantizeActivationNvFp4Kernel<<>>( + reinterpret_cast(a), input_scale, reinterpret_cast(a_packed), + reinterpret_cast(a_scale), alpha, weight_scale_2, m, k, rounded_k_blocks); + } + ORT_RETURN_IF_ERROR(CUDA_CALL(cudaGetLastError())); + + if (is_bf16) { + return DispatchRunGemm( + y, a_packed, b_packed, a_scale, b_scale, alpha, m, n, k, workspace, stream); + } + return DispatchRunGemm( + y, a_packed, b_packed, a_scale, b_scale, alpha, m, n, k, workspace, stream); +} + +} // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index eb1c1edd378c1..8013429fd7604 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -3005,6 +3005,63 @@ ONNX_MS_OPERATOR_SET_SCHEMA(GemmFloat8, 1, updateOutputShape(ctx, 0, {first_input_shape.dim(transA ? 1 : 0), second_input_shape.dim(transB ? 0 : 1)}); })); +ONNX_MS_OPERATOR_SET_SCHEMA( + MatMulBlockScaledFp4, 1, + OpSchema() + .SetDoc(R"DOC(Weight-only NVFP4 (E2M1) matrix multiplication. + +The weight tensor B is stored as packed NVFP4: two E2M1 values per byte (low nibble first). +The dequantized weight value is `e2m1(B) * weight_scale_2 * e4m3(weight_scale[n, k / block_size])`, +where `weight_scale` holds one E4M3 scale per `block_size` (default 16) consecutive K values and +`weight_scale_2` is a single global fp32 scale. The weight is dequantized to the activation type +(FP16/BF16) and multiplied with the FP16/BF16 activation. This path is architecture independent and +runs on Hopper (SM90) as well as Blackwell.)DOC") + .Attr("K", "Inner (contraction) dimension: the number of logical columns of the unpacked weight.", + AttributeProto::INT) + .Attr("N", "Number of output columns, i.e. the number of rows of the packed weight.", + AttributeProto::INT) + .Attr("block_size", "Number of consecutive K values that share one E4M3 weight scale. Default 16.", + AttributeProto::INT, static_cast(16)) + .Input(0, "A", "Row-major FP16/BF16 activation of shape [..., K].", "T") + .Input(1, "B", + "Packed NVFP4 weight of shape [N, K/2] stored as uint8 (two E2M1 values per byte, low nibble first).", + "T1") + .Input(2, "weight_scale", + "Per-block E4M3 weight scales of shape [N, ceil(K / block_size)] stored as raw uint8 bytes.", "T2") + .Input(3, "weight_scale_2", "Global fp32 weight scale (scalar).", "T3") + .Input(4, "input_scale", + "Optional global fp32 activation scale (scalar). Accepted for parity with quantized checkpoints; " + "it is a no-op on the weight-only FP16/BF16 path and is reserved for the native NVFP4 path on Blackwell.", + "T3", OpSchema::Optional) + .Input(5, "bias", "Optional bias of shape [N].", "T", OpSchema::Optional) + .Output(0, "Y", "Output of shape [..., N] in the activation type.", "T") + .TypeConstraint("T", {"tensor(float16)", "tensor(bfloat16)"}, + "Constrain activation, bias and output to FP16 or BF16.") + .TypeConstraint("T1", {"tensor(uint8)"}, "Constrain packed NVFP4 weight to uint8.") + .TypeConstraint("T2", {"tensor(uint8)"}, "Constrain E4M3 weight scales to uint8.") + .TypeConstraint("T3", {"tensor(float)"}, "Constrain scalar scales to FP32.") + .TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) { + propagateElemTypeFromInputToOutput(ctx, 0, 0); + if (!hasInputShape(ctx, 0)) { + return; + } + const auto& a_shape = getInputShape(ctx, 0); + if (a_shape.dim_size() < 1) { + fail_shape_inference("A must have rank at least 1."); + } + ONNX_NAMESPACE::TensorShapeProto output_shape; + for (int i = 0; i < a_shape.dim_size() - 1; ++i) { + *output_shape.add_dim() = a_shape.dim(i); + } + const auto* n_attr = ctx.getAttribute("N"); + if (n_attr != nullptr && n_attr->has_i()) { + output_shape.add_dim()->set_dim_value(n_attr->i()); + } else { + output_shape.add_dim(); + } + updateOutputShape(ctx, 0, output_shape); + })); + static void MatmulWithQuantWeightShapeInference(ONNX_NAMESPACE::InferenceContext& ctx, int64_t K, int64_t N, diff --git a/onnxruntime/core/graph/contrib_ops/ms_opset.h b/onnxruntime/core/graph/contrib_ops/ms_opset.h index 59f97c222ceb2..fbc259575708f 100644 --- a/onnxruntime/core/graph/contrib_ops/ms_opset.h +++ b/onnxruntime/core/graph/contrib_ops/ms_opset.h @@ -121,6 +121,7 @@ class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GemmFastGelu); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, DecoderMaskedSelfAttention); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, DecoderMaskedMultiHeadAttention); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GemmFloat8); +class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulBlockScaledFp4); class OpSet_Microsoft_ver1 { public: @@ -236,6 +237,7 @@ class OpSet_Microsoft_ver1 { fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); + fn(GetOpSchema()); } }; } // namespace contrib diff --git a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc new file mode 100644 index 0000000000000..7e1c0c03172d7 --- /dev/null +++ b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc @@ -0,0 +1,207 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "gtest/gtest.h" +#include "test/common/cuda_op_test_utils.h" +#include "test/common/tensor_op_test_utils.h" +#include "test/providers/provider_test_utils.h" +#include "test/unittest_util/conversion.h" + +namespace onnxruntime::test { + +#if defined(USE_CUDA) + +// NVFP4 (E2M1) 4-bit magnitude nibble encodings (sign bit is 0x8): +// +0.0 -> 0x0, +0.5 -> 0x1, +1.0 -> 0x2, +1.5 -> 0x3, +// +2.0 -> 0x4, +3.0 -> 0x5, +4.0 -> 0x6, +6.0 -> 0x7 +// A packed byte holds two values: low nibble is element 2j, high nibble is element 2j+1. +// +// E4M3 (float8e4m3fn) scale byte encodings: +// 1.0 -> 0x38, 2.0 -> 0x40, 0.5 -> 0x30 + +// A -> [M, K] all ones per row scaled by (m + 1); weights are constant per row, so the +// operator must reproduce Y[m, n] = W_val[n] * sum_k A[m, k]. +TEST(MatMulBlockScaledFp4OpTest, WeightOnlyBasicFp16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + } + + constexpr int64_t m = 2; + constexpr int64_t n = 2; + constexpr int64_t k = 16; // one block with block_size == 16 + + // Weight row 0 = +1.0 (nibble 0x2 -> byte 0x22), row 1 = +2.0 (nibble 0x4 -> byte 0x44). + std::vector b(n * (k / 2)); + for (int64_t j = 0; j < k / 2; ++j) { + b[0 * (k / 2) + j] = 0x22; + b[1 * (k / 2) + j] = 0x44; + } + // One E4M3 scale per row (single K block), both 1.0. + std::vector weight_scale = {0x38, 0x38}; + std::vector weight_scale_2 = {1.0f}; + + std::vector a(m * k); + for (int64_t row = 0; row < m; ++row) { + for (int64_t col = 0; col < k; ++col) { + a[row * k + col] = static_cast(row + 1); + } + } + // W[0, :] = 1.0, W[1, :] = 2.0; sum_k A[0, :] = 16, sum_k A[1, :] = 32. + std::vector expected = {16.0f, 32.0f, 32.0f, 64.0f}; + + OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + test.AddAttribute("K", k); + test.AddAttribute("N", n); + test.AddAttribute("block_size", 16); + test.AddInput("A", {m, k}, FloatsToMLFloat16s(a)); + test.AddInput("B", {n, k / 2}, b); + test.AddInput("weight_scale", {n, 1}, weight_scale); + test.AddInput("weight_scale_2", {1}, weight_scale_2); + test.AddOutput("Y", {m, n}, FloatsToMLFloat16s(expected)); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +// Exercises non-unit per-block E4M3 scales, a global weight_scale_2, negative weights, bias and +// a skipped optional input_scale, with BF16 activations/output. +TEST(MatMulBlockScaledFp4OpTest, WeightOnlyScalesBiasBf16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + } + + constexpr int64_t m = 1; + constexpr int64_t n = 2; + constexpr int64_t k = 16; + + // Weight row 0 = +1.0 (0x22), row 1 = -1.0 (nibble 0xA -> byte 0xAA). + std::vector b(n * (k / 2)); + for (int64_t j = 0; j < k / 2; ++j) { + b[0 * (k / 2) + j] = 0x22; + b[1 * (k / 2) + j] = 0xAA; + } + // Row 0 scale = 2.0 (0x40), row 1 scale = 1.0 (0x38). + std::vector weight_scale = {0x40, 0x38}; + std::vector weight_scale_2 = {3.0f}; + + std::vector a(m * k, 1.0f); + // W[0, :] = 1.0 * 3.0 * 2.0 = 6.0; W[1, :] = -1.0 * 3.0 * 1.0 = -3.0; sum_k A = 16. + // Y = {6*16, -3*16} + bias{1, 2} = {97, -46}. + std::vector bias = {1.0f, 2.0f}; + std::vector expected = {97.0f, -46.0f}; + + OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + test.AddAttribute("K", k); + test.AddAttribute("N", n); + test.AddAttribute("block_size", 16); + test.AddInput("A", {m, k}, FloatsToBFloat16s(a)); + test.AddInput("B", {n, k / 2}, b); + test.AddInput("weight_scale", {n, 1}, weight_scale); + test.AddInput("weight_scale_2", {1}, weight_scale_2); + test.AddOptionalInputEdge(); // input_scale (skipped) + test.AddInput("bias", {n}, FloatsToBFloat16s(bias)); + test.AddOutput("Y", {m, n}, FloatsToBFloat16s(expected)); + test.SetOutputTolerance(0.5f); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +// Exercises the fused decode GEMV fast path (small M) with a multi-block K (K = 64 == 4 blocks, +// K % 32 == 0), FP16 activations. Weights are constant per row so Y[m, n] = W_val[n] * sum_k A[m, k]. +TEST(MatMulBlockScaledFp4OpTest, GemvDecodeMultiBlockFp16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + } + + constexpr int64_t m = 2; + constexpr int64_t n = 2; + constexpr int64_t k = 64; // 4 blocks with block_size == 16, K % 32 == 0 -> GEMV path + constexpr int64_t k_blocks = k / 16; + + // Weight row 0 = +1.0 (byte 0x22), row 1 = +2.0 (byte 0x44). + std::vector b(n * (k / 2)); + for (int64_t j = 0; j < k / 2; ++j) { + b[0 * (k / 2) + j] = 0x22; + b[1 * (k / 2) + j] = 0x44; + } + // Unit E4M3 scale (0x38 == 1.0) for every block of every row. + std::vector weight_scale(n * k_blocks, 0x38); + std::vector weight_scale_2 = {1.0f}; + + std::vector a(m * k); + for (int64_t row = 0; row < m; ++row) { + for (int64_t col = 0; col < k; ++col) { + a[row * k + col] = static_cast(row + 1); + } + } + // W[0, :] = 1.0, W[1, :] = 2.0; sum_k A[0, :] = 64, sum_k A[1, :] = 128. + std::vector expected = {64.0f, 128.0f, 128.0f, 256.0f}; + + OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + test.AddAttribute("K", k); + test.AddAttribute("N", n); + test.AddAttribute("block_size", 16); + test.AddInput("A", {m, k}, FloatsToMLFloat16s(a)); + test.AddInput("B", {n, k / 2}, b); + test.AddInput("weight_scale", {n, k_blocks}, weight_scale); + test.AddInput("weight_scale_2", {1}, weight_scale_2); + test.AddOutput("Y", {m, n}, FloatsToMLFloat16s(expected)); + test.SetOutputTolerance(0.5f); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +// Exercises the fused decode GEMV fast path with M == 1, per-block scales, a global weight_scale_2, +// negative weights and bias (BF16). K = 32 == 2 blocks, K % 32 == 0 -> GEMV path. +TEST(MatMulBlockScaledFp4OpTest, GemvDecodeScalesBiasBf16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + } + + constexpr int64_t m = 1; + constexpr int64_t n = 2; + constexpr int64_t k = 32; // 2 blocks with block_size == 16 + constexpr int64_t k_blocks = k / 16; + + // Weight row 0 = +1.0 (0x22), row 1 = -1.0 (nibble 0xA -> byte 0xAA). + std::vector b(n * (k / 2)); + for (int64_t j = 0; j < k / 2; ++j) { + b[0 * (k / 2) + j] = 0x22; + b[1 * (k / 2) + j] = 0xAA; + } + // Row 0 scale = 2.0 (0x40) for both blocks, row 1 scale = 1.0 (0x38) for both blocks. + std::vector weight_scale = {0x40, 0x40, 0x38, 0x38}; + std::vector weight_scale_2 = {3.0f}; + + std::vector a(m * k, 1.0f); + // W[0, :] = 1.0 * 3.0 * 2.0 = 6.0; W[1, :] = -1.0 * 3.0 * 1.0 = -3.0; sum_k A = 32. + // Y = {6*32, -3*32} + bias{1, 2} = {193, -94}. + std::vector bias = {1.0f, 2.0f}; + std::vector expected = {193.0f, -94.0f}; + + OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + test.AddAttribute("K", k); + test.AddAttribute("N", n); + test.AddAttribute("block_size", 16); + test.AddInput("A", {m, k}, FloatsToBFloat16s(a)); + test.AddInput("B", {n, k / 2}, b); + test.AddInput("weight_scale", {n, k_blocks}, weight_scale); + test.AddInput("weight_scale_2", {1}, weight_scale_2); + test.AddOptionalInputEdge(); // input_scale (skipped) + test.AddInput("bias", {n}, FloatsToBFloat16s(bias)); + test.AddOutput("Y", {m, n}, FloatsToBFloat16s(expected)); + test.SetOutputTolerance(0.5f); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} + +#endif // USE_CUDA + +} // namespace onnxruntime::test diff --git a/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py b/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py new file mode 100644 index 0000000000000..b2c1644b924b0 --- /dev/null +++ b/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py @@ -0,0 +1,398 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +""" +Accuracy and latency harness for the CUDA MatMulBlockScaledFp4 contrib op. + +The script builds a single-node com.microsoft contrib-op model, binds CUDA tensors with +I/O binding, compares the output with an FP32 dequantized reference, and prints one JSON +record per case. It is intended for opt-in Blackwell profiling, not for normal CI. + +Examples: + python profile_matmul_block_scaled.py --suite smoke + python profile_matmul_block_scaled.py --op fp4 --activation-dtype bf16 --m 1 --n 4096 --k 4096 --bias + python profile_matmul_block_scaled.py --op fp4 --m 16 --n 11008 --k 4096 --repeat 200 + +For kernel-level evidence, wrap a representative case with nsys: + nsys profile -t cuda,nvtx -o block_scaled --export=sqlite \ + python profile_matmul_block_scaled.py --op fp4 --m 16 --n 4096 --k 4096 +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import statistics +import time +from contextlib import nullcontext +from dataclasses import dataclass +from typing import Any + +import numpy as np +import torch +from onnx import TensorProto, helper + +import onnxruntime +from onnxruntime.capi.onnxruntime_pybind11_state import Fail as OrtFail + +try: + import nvtx + + _HAS_NVTX = True +except ImportError: + nvtx = None + _HAS_NVTX = False + + +RESULT_PREFIX = "MATMUL_BLOCK_SCALED_RESULT " + +_TORCH_TO_ONNX = { + torch.float16: TensorProto.FLOAT16, + torch.bfloat16: TensorProto.BFLOAT16, +} +_FP4_POS_VALUES = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32) + + +@dataclass(frozen=True) +class Case: + op: str + m: int + n: int + k: int + activation_dtype: str + block_size: int | None = None + bias: bool = False + seed: int = 0 + + +def _nvtx_range(name: str, color: str = "green"): + if not _HAS_NVTX: + return nullcontext() + return nvtx.annotate(name, color=color) + + +def _require_cuda() -> None: + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required for this harness.") + if "CUDAExecutionProvider" not in onnxruntime.get_available_providers(): + raise RuntimeError("CUDAExecutionProvider is not available in this onnxruntime build.") + + +def _torch_dtype(name: str) -> torch.dtype: + if name == "fp16": + return torch.float16 + if name == "bf16": + return torch.bfloat16 + raise ValueError(f"Unsupported dtype: {name}") + + +def _onnx_dtype(name: str) -> int: + return _TORCH_TO_ONNX[_torch_dtype(name)] + + +def _raw_uint8(tensor: torch.Tensor) -> bytes: + return np.ascontiguousarray(tensor.detach().view(torch.uint8).cpu().numpy()).tobytes() + + +def _make_float_initializer(name: str, tensor: torch.Tensor, onnx_dtype: int): + if onnx_dtype == TensorProto.FLOAT: + values = np.ascontiguousarray(tensor.detach().cpu().numpy().astype(np.float32)) + return helper.make_tensor(name, onnx_dtype, list(tensor.shape), values.tobytes(), raw=True) + if onnx_dtype == TensorProto.FLOAT16: + values = np.ascontiguousarray(tensor.detach().cpu().numpy().astype(np.float16)) + return helper.make_tensor(name, onnx_dtype, list(tensor.shape), values.tobytes(), raw=True) + if onnx_dtype == TensorProto.BFLOAT16: + values = tensor.detach().to(torch.float32).flatten().cpu().tolist() + return helper.make_tensor(name, onnx_dtype, list(tensor.shape), values, raw=False) + raise ValueError(f"Unsupported initializer dtype: {onnx_dtype}") + + +def _make_session(model: bytes) -> onnxruntime.InferenceSession: + session_options = onnxruntime.SessionOptions() + session_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL + session_options.log_severity_level = 3 + try: + return onnxruntime.InferenceSession(model, session_options, providers=["CUDAExecutionProvider"]) + except OrtFail as error: + if "MatMulBlockScaled" in str(error) and "not a registered" in str(error): + raise RuntimeError( + "The active onnxruntime package does not register the MatMulBlockScaled contrib ops. " + "Build and install this branch's CUDA wheel before running the harness." + ) from error + raise + + +def _model_bytes(nodes, graph_inputs, graph_outputs, initializers, name: str) -> bytes: + graph = helper.make_graph(nodes, name, graph_inputs, graph_outputs, initializers) + model = helper.make_model( + graph, + opset_imports=[helper.make_opsetid("com.microsoft", 1), helper.make_opsetid("", 17)], + ) + return model.SerializeToString() + + +def _quantize_fp4(weight: torch.Tensor, block_size: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + if weight.shape[1] % 2 != 0: + raise ValueError("FP4 packed weight requires even K.") + + n, k = weight.shape + k_blocks = math.ceil(k / block_size) + padded_k = k_blocks * block_size + padded = torch.nn.functional.pad(weight.float(), (0, padded_k - k)) + blocks = padded.reshape(n, k_blocks, block_size) + max_abs = blocks.abs().amax(dim=-1) + scale = torch.clamp(max_abs / 6.0, min=1.0 / 1024.0).to(torch.float8_e4m3fn) + scale_f32 = scale.float() + + scaled = blocks / scale_f32.unsqueeze(-1) + values = _FP4_POS_VALUES.to(device=weight.device) + flat_abs = scaled.abs().reshape(-1, 1).clamp(max=6.0) + nearest = torch.abs(flat_abs - values.reshape(1, -1)).argmin(dim=1).reshape_as(scaled).to(torch.uint8) + codes = nearest | ((scaled < 0).to(torch.uint8) << 3) + codes = codes.reshape(n, padded_k)[:, :k].contiguous() + + low = codes[:, 0::2] + high = codes[:, 1::2] + packed = (low | (high << 4)).contiguous() + + quantized_values = values[nearest.long()].reshape_as(scaled) * torch.where(scaled < 0, -1.0, 1.0) + dequantized = (quantized_values * scale_f32.unsqueeze(-1)).reshape(n, padded_k)[:, :k].contiguous() + return packed, scale.view(torch.uint8).contiguous(), dequantized + + +def _make_fp4_model(case: Case, b_packed: torch.Tensor, weight_scale: torch.Tensor, bias: torch.Tensor | None) -> bytes: + activation_onnx_type = _onnx_dtype(case.activation_dtype) + block_size = case.block_size or 16 + inputs = ["A", "B", "weight_scale", "weight_scale_2"] + initializers = [ + helper.make_tensor("B", TensorProto.UINT8, [case.n, case.k // 2], _raw_uint8(b_packed), raw=True), + helper.make_tensor( + "weight_scale", + TensorProto.UINT8, + [case.n, math.ceil(case.k / block_size)], + _raw_uint8(weight_scale), + raw=True, + ), + helper.make_tensor("weight_scale_2", TensorProto.FLOAT, [1], [1.0], raw=False), + ] + if bias is not None: + inputs.extend(["", "bias"]) + initializers.append(_make_float_initializer("bias", bias, activation_onnx_type)) + + node = helper.make_node( + "MatMulBlockScaledFp4", + inputs, + ["Y"], + domain="com.microsoft", + K=case.k, + N=case.n, + block_size=block_size, + ) + graph_inputs = [helper.make_tensor_value_info("A", activation_onnx_type, [case.m, case.k])] + graph_outputs = [helper.make_tensor_value_info("Y", activation_onnx_type, [case.m, case.n])] + return _model_bytes([node], graph_inputs, graph_outputs, initializers, "MatMulBlockScaledFp4_Profile") + + +def _fp4_reference(a: torch.Tensor, b_dequantized: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor: + result = a.float() @ b_dequantized.float().T + if bias is not None: + result += bias.float().reshape(1, -1) + return result.to(a.dtype).float() + + +def _fp4_native_sm120_enabled() -> bool: + return os.environ.get("ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120", "").lower() in {"1", "true", "yes", "on"} + + +def _fp4_native_sm120_supported(case: Case) -> bool: + block_size = case.block_size or 16 + return ( + _fp4_native_sm120_enabled() + and case.m > 8 + and block_size == 16 + and case.k % 32 == 0 + and case.n % 32 == 0 + and case.activation_dtype in {"fp16", "bf16"} + ) + + +def _fp4_expected_path(case: Case) -> str: + block_size = case.block_size or 16 + if case.m > 0 and case.m <= 8 and block_size == 16 and case.k % 32 == 0: + return "fp4_gemv" + if _fp4_native_sm120_supported(case): + return "sm120_native_fp4_gemm" + return "fp4_dequant_cublas" + + +def _make_inputs(case: Case) -> tuple[bytes, torch.Tensor, torch.Tensor, str]: + generator = torch.Generator(device="cuda") + generator.manual_seed(case.seed) + block_size = case.block_size or 16 + + activation_dtype = _torch_dtype(case.activation_dtype) + a = (torch.randn((case.m, case.k), generator=generator, device="cuda") * 0.75).to(activation_dtype).contiguous() + weight = torch.randn((case.n, case.k), generator=generator, device="cuda", dtype=torch.float32) * 0.75 + b_packed, weight_scale, b_dequantized = _quantize_fp4(weight, block_size) + bias = None + if case.bias: + bias = (torch.randn((case.n,), generator=generator, device="cuda") * 0.25).to(activation_dtype).contiguous() + model = _make_fp4_model(case, b_packed, weight_scale, bias) + if _fp4_native_sm120_supported(case): + _, _, a_dequantized = _quantize_fp4(a.float(), block_size) + reference = _fp4_reference(a_dequantized.to(activation_dtype), b_dequantized, bias) + else: + reference = _fp4_reference(a, b_dequantized, bias) + return model, a, reference, _fp4_expected_path(case) + + +def _error_metrics(actual: torch.Tensor, expected: torch.Tensor) -> dict[str, float]: + diff = actual.float() - expected.float() + abs_diff = diff.abs() + rel_diff = abs_diff / torch.clamp(expected.float().abs(), min=1.0e-6) + return { + "max_abs_error": float(abs_diff.max().item()) if abs_diff.numel() else 0.0, + "max_rel_error": float(rel_diff.max().item()) if rel_diff.numel() else 0.0, + "rmse": float(torch.sqrt(torch.mean(diff * diff)).item()) if diff.numel() else 0.0, + "max_expected_abs": float(expected.float().abs().max().item()) if expected.numel() else 0.0, + } + + +def _run_timed( + session: onnxruntime.InferenceSession, a: torch.Tensor, y: torch.Tensor, warmup: int, repeat: int +) -> list[float]: + io_binding = session.io_binding() + io_binding.bind_input("A", "cuda", 0, _TORCH_TO_ONNX[a.dtype], list(a.shape), a.data_ptr()) + io_binding.bind_output("Y", "cuda", 0, _TORCH_TO_ONNX[y.dtype], list(y.shape), y.data_ptr()) + + with _nvtx_range("warmup", "yellow"): + for _ in range(warmup): + session.run_with_iobinding(io_binding) + torch.cuda.synchronize() + + times_ms = [] + with _nvtx_range("benchmark", "green"): + for _ in range(repeat): + start = time.perf_counter() + session.run_with_iobinding(io_binding) + torch.cuda.synchronize() + times_ms.append((time.perf_counter() - start) * 1000.0) + return times_ms + + +def _summarize_times(times_ms: list[float]) -> dict[str, float]: + sorted_times = sorted(times_ms) + p90_index = min(len(sorted_times) - 1, math.ceil(0.90 * len(sorted_times)) - 1) + p99_index = min(len(sorted_times) - 1, math.ceil(0.99 * len(sorted_times)) - 1) + return { + "mean_ms": statistics.fmean(times_ms), + "p50_ms": statistics.median(times_ms), + "p90_ms": sorted_times[p90_index], + "p99_ms": sorted_times[p99_index], + "min_ms": sorted_times[0], + } + + +def run_case(case: Case, warmup: int, repeat: int, atol: float, rtol: float) -> dict[str, Any]: + model, a, reference, expected_path = _make_inputs(case) + output_dtype = _torch_dtype(case.activation_dtype) + y = torch.empty((case.m, case.n), dtype=output_dtype, device="cuda") + session = _make_session(model) + times_ms = _run_timed(session, a, y, warmup, repeat) + + metrics = _error_metrics(y, reference) + threshold = atol + rtol * metrics["max_expected_abs"] + passed = metrics["max_abs_error"] <= threshold + flops = 2.0 * case.m * case.n * case.k + timing = _summarize_times(times_ms) + result = { + "op": case.op, + "m": case.m, + "n": case.n, + "k": case.k, + "block_size": case.block_size or 16, + "activation_dtype": case.activation_dtype, + "bias": case.bias, + "expected_path": expected_path, + "passed": passed, + "atol": atol, + "rtol": rtol, + "tflops": flops / (timing["mean_ms"] * 1.0e-3) / 1.0e12, + **timing, + **metrics, + } + print(RESULT_PREFIX + json.dumps(result, sort_keys=True)) + return result + + +def _default_cases(args) -> list[Case]: + if args.m is not None and args.n is not None and args.k is not None: + return [ + Case( + op=args.op, + m=args.m, + n=args.n, + k=args.k, + activation_dtype=args.activation_dtype, + block_size=args.block_size, + bias=args.bias, + seed=args.seed, + ) + ] + + cases = [] + if args.suite == "smoke": + cases.extend( + [ + Case("fp4", 1, 80, 256, "fp16", bias=True, seed=args.seed + 3), + Case("fp4", 32, 128, 256, "bf16", bias=False, seed=args.seed + 4), + ] + ) + return cases + + matrix_ms = [1, 2, 4, 8] if args.suite == "decode" else [16, 32, 64, 128] + matrix_shapes = [(4096, 4096), (4096, 11008)] + for k, n in matrix_shapes: + cases.extend( + Case("fp4", m, n, k, args.activation_dtype, bias=args.bias, seed=args.seed + 100 + m) for m in matrix_ms + ) + return cases + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Profile the CUDA block-scaled FP4 MatMul contrib op") + parser.add_argument("--op", choices=["fp4"], default="fp4") + parser.add_argument("--suite", choices=["smoke", "decode", "prefill"], default="smoke") + parser.add_argument("--m", type=int, help="M rows for single-case mode") + parser.add_argument("--n", type=int, help="N columns for single-case mode") + parser.add_argument("--k", type=int, help="K reduction dimension for single-case mode") + parser.add_argument("--block-size", type=int, help="Override block_size attribute") + parser.add_argument("--activation-dtype", choices=["fp16", "bf16"], default="fp16") + parser.add_argument("--bias", action="store_true", help="Enable FP4 bias") + parser.add_argument("--warmup", type=int, default=10) + parser.add_argument("--repeat", type=int, default=50) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--atol", type=float, default=2.0) + parser.add_argument("--rtol", type=float, default=0.02) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + _require_cuda() + single_case_args = [args.m is not None, args.n is not None, args.k is not None] + if any(single_case_args) and not all(single_case_args): + raise ValueError("Single-case mode requires all of --m, --n and --k.") + + results = [run_case(case, args.warmup, args.repeat, args.atol, args.rtol) for case in _default_cases(args)] + failures = [result for result in results if not result["passed"]] + if failures: + raise SystemExit(f"{len(failures)} case(s) failed accuracy checks") + + +if __name__ == "__main__": + main() From 2ee990aa9b48f1bfa1e890f640a82cd0f7349437 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Wed, 22 Jul 2026 00:42:38 -0700 Subject: [PATCH 02/11] Address PR feedback: portable UseM256Config, strict shape validation, GEMV guards, doc env vars - Replace __builtin_clz with a portable next-power-of-two loop so MSVC host builds compile. - Validate exact ranks/dims of B, weight_scale, and bias in ComputeImpl instead of only total element counts. - Enforce rank-2 [N, K/16] weight_scale shape in PrePack. - Add block_size==16 and K%32==0 runtime guards in LaunchMatMulBlockScaledFp4Gemv. - Replace hard-coded absolute paths in docs with ORT_REPO/ORT_BUILD variables. --- .../cuda/matmul_block_scaled_fp4.md | 31 +++++++++++------- .../matmul_block_scaled_fp4_experiments.md | 32 ++++++++++++------- .../cuda/math/matmul_block_scaled_fp4.cc | 21 ++++++++---- .../cuda/math/matmul_block_scaled_fp4.cu | 5 +++ .../math/matmul_block_scaled_fp4_sm120.cu | 8 +++-- 5 files changed, 65 insertions(+), 32 deletions(-) diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md index 14e2a03fcc9f1..07413932d266e 100644 --- a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md @@ -200,10 +200,19 @@ The default remains the existing weight-only semantics: decode GEMV for small ## 9. Testing and Benchmarking +The commands below use two environment variables so they can be copied without +editing developer-specific paths. Set them once to your repo root and build +output directory: + +```bash +export ORT_REPO=$(git rev-parse --show-toplevel) +export ORT_BUILD="$ORT_REPO/build/cu130/Release" +``` + Focused C++ tests: ```bash -CUDA_VISIBLE_DEVICES=0 build/cu130/Release/onnxruntime_provider_test \ +CUDA_VISIBLE_DEVICES=0 "$ORT_BUILD/onnxruntime_provider_test" \ --gtest_filter='MatMulBlockScaledFp4OpTest.*' ``` @@ -211,19 +220,19 @@ Python harness examples: ```bash # Decode GEMV -cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ - python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ +cd /tmp && PYTHONPATH="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ + python "$ORT_REPO/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py" \ --op fp4 --activation-dtype fp16 --m 1 --n 11008 --k 4096 --warmup 100 --repeat 500 # Default prefill: dequantize + cuBLAS -cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ - python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ +cd /tmp && PYTHONPATH="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ + python "$ORT_REPO/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py" \ --op fp4 --activation-dtype fp16 --m 16 --n 11008 --k 4096 --warmup 50 --repeat 200 # Native SM120 prefill -cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ +cd /tmp && PYTHONPATH="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ - python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + python "$ORT_REPO/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py" \ --op fp4 --activation-dtype fp16 --m 16 --n 11008 --k 4096 --warmup 50 --repeat 200 ``` @@ -231,8 +240,8 @@ After rebuilding `libonnxruntime_providers_cuda.so`, sync the provider into the Python load locations before Python benchmarks: ```bash -cp build/cu130/Release/libonnxruntime_providers_cuda.so \ - build/cu130/Release/onnxruntime/capi/libonnxruntime_providers_cuda.so -cp build/cu130/Release/libonnxruntime_providers_cuda.so \ - build/cu130/Release/build/lib/onnxruntime/capi/libonnxruntime_providers_cuda.so +cp "$ORT_BUILD/libonnxruntime_providers_cuda.so" \ + "$ORT_BUILD/onnxruntime/capi/libonnxruntime_providers_cuda.so" +cp "$ORT_BUILD/libonnxruntime_providers_cuda.so" \ + "$ORT_BUILD/build/lib/onnxruntime/capi/libonnxruntime_providers_cuda.so" ``` diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md index e940f8ba91dd0..f70867e1b40b1 100644 --- a/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md @@ -128,43 +128,51 @@ unswizzled scale layout. ## 4. Benchmark Commands +The commands below use `ORT_REPO` and `ORT_BUILD` so they can be copied without +editing developer-specific paths. Set them once: + +```bash +export ORT_REPO=$(git rev-parse --show-toplevel) +export ORT_BUILD="$ORT_REPO/build/cu130/Release" +``` + Provider rebuild and Python-provider sync: ```bash -cmake --build build/cu130/Release --target onnxruntime_providers_cuda --parallel -cp build/cu130/Release/libonnxruntime_providers_cuda.so \ - build/cu130/Release/onnxruntime/capi/libonnxruntime_providers_cuda.so -cp build/cu130/Release/libonnxruntime_providers_cuda.so \ - build/cu130/Release/build/lib/onnxruntime/capi/libonnxruntime_providers_cuda.so +cmake --build "$ORT_BUILD" --target onnxruntime_providers_cuda --parallel +cp "$ORT_BUILD/libonnxruntime_providers_cuda.so" \ + "$ORT_BUILD/onnxruntime/capi/libonnxruntime_providers_cuda.so" +cp "$ORT_BUILD/libonnxruntime_providers_cuda.so" \ + "$ORT_BUILD/build/lib/onnxruntime/capi/libonnxruntime_providers_cuda.so" ``` Decode benchmarks: ```bash -cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ +cd /tmp && PYTHONPATH="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ - python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + python "$ORT_REPO/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py" \ --op fp4 --activation-dtype fp16 --m 1 --n 11008 --k 4096 --warmup 100 --repeat 500 -cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ +cd /tmp && PYTHONPATH="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ - python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + python "$ORT_REPO/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py" \ --op fp4 --activation-dtype fp16 --m 8 --n 11008 --k 4096 --warmup 100 --repeat 500 ``` Native prefill benchmark: ```bash -cd /tmp && PYTHONPATH=/home/tlwu/onnxruntime/build/cu130/Release CUDA_VISIBLE_DEVICES=0 \ +cd /tmp && PYTHONPATH="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ - python /home/tlwu/onnxruntime/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py \ + python "$ORT_REPO/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py" \ --op fp4 --activation-dtype fp16 --m 16 --n 11008 --k 4096 --warmup 50 --repeat 200 ``` Focused C++ tests: ```bash -CUDA_VISIBLE_DEVICES=0 build/cu130/Release/onnxruntime_provider_test \ +CUDA_VISIBLE_DEVICES=0 "$ORT_BUILD/onnxruntime_provider_test" \ --gtest_filter='MatMulBlockScaledFp4OpTest.*' ``` diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index 45f35788cea24..38da1b064715b 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -63,8 +63,10 @@ Status MatMulBlockScaledFp4::PrePack(const Tensor& tensor, int input_idx, Alloca } const int64_t k_blocks = K_ / 16; - ORT_RETURN_IF_NOT(tensor.Shape().Size() >= N_ * k_blocks, - "weight_scale tensor is too small; expected at least ", N_ * k_blocks, " E4M3 scales."); + const auto& scale_shape = tensor.Shape(); + ORT_RETURN_IF_NOT(scale_shape.NumDimensions() == 2 && scale_shape[0] == N_ && scale_shape[1] == k_blocks, + "weight_scale must have shape [N, K/16] = [", N_, ", ", k_blocks, "], got ", + scale_shape.ToString(), "."); const int64_t rounded_k_blocks = RoundUp(k_blocks, 4); const int64_t rounded_n = RoundUp(N_, 128); @@ -112,10 +114,14 @@ Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { const int64_t k_packed = K_ / 2; const int64_t k_blocks = (K_ + block_size_ - 1) / block_size_; - ORT_ENFORCE(b->Shape().Size() >= N_ * k_packed, - "B tensor is too small; expected at least ", N_ * k_packed, " packed bytes."); - ORT_ENFORCE(weight_scale->Shape().Size() >= N_ * k_blocks, - "weight_scale tensor is too small; expected at least ", N_ * k_blocks, " E4M3 scales."); + const auto& b_shape = b->Shape(); + ORT_ENFORCE(b_shape.NumDimensions() == 2 && b_shape[0] == N_ && b_shape[1] == k_packed, + "B must have shape [N, K/2] = [", N_, ", ", k_packed, "], got ", b_shape.ToString(), "."); + const auto& weight_scale_shape = weight_scale->Shape(); + ORT_ENFORCE(weight_scale_shape.NumDimensions() == 2 && weight_scale_shape[0] == N_ && + weight_scale_shape[1] == k_blocks, + "weight_scale must have shape [N, ceil(K/block_size)] = [", N_, ", ", k_blocks, "], got ", + weight_scale_shape.ToString(), "."); ORT_ENFORCE(weight_scale_2->Shape().Size() == 1, "weight_scale_2 must be a scalar."); if (input_scale != nullptr) { ORT_ENFORCE(input_scale->Shape().Size() == 1, "input_scale must be a scalar."); @@ -123,7 +129,8 @@ Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { // weight-only FP16/BF16 activation path keeps full-precision activations. } if (bias != nullptr) { - ORT_ENFORCE(bias->Shape().Size() == N_, "bias must have shape [N]."); + ORT_ENFORCE(bias->Shape().NumDimensions() == 1 && bias->Shape()[0] == N_, + "bias must have shape [N] = [", N_, "], got ", bias->Shape().ToString(), "."); } constexpr bool transa = false; diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu index 26e694e6ce0a0..a544a6d943530 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu @@ -312,6 +312,11 @@ Status LaunchMatMulBlockScaledFp4Gemv(void* y, if (m <= 0 || n <= 0 || k <= 0) { return Status::OK(); } + // This kernel is hard-coded for block_size == 16 and assumes K is a multiple of 32 so that each + // warp lane always owns a full 32-element slice (and one E4M3 scale per 16-element block). Guard + // against misuse if this helper is ever reused outside the callers that already check these. + ORT_RETURN_IF_NOT(block_size == 16, "MatMulBlockScaledFp4 GEMV requires block_size == 16, got ", block_size, "."); + ORT_RETURN_IF_NOT(k % 32 == 0, "MatMulBlockScaledFp4 GEMV requires K divisible by 32, got ", k, "."); const int k_blocks = (k + block_size - 1) / block_size; constexpr int kWarpsPerBlock = 8; const dim3 threads{32, kWarpsPerBlock}; diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu index 9ed24f97a39e2..59911bfe3f0e1 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu @@ -248,8 +248,12 @@ Status RunGemm(void* y, } bool UseM256Config(int m) { - const auto m_unsigned = static_cast(std::max(m - 1, 1)); - const int next_power_of_two_m = static_cast(1u << (32 - __builtin_clz(m_unsigned))); + // Smallest power of two >= m, computed portably (no compiler builtins so this + // also compiles under MSVC host compilation for CUDA builds on Windows). + int next_power_of_two_m = 1; + while (next_power_of_two_m < m) { + next_power_of_two_m <<= 1; + } return std::max(16, next_power_of_two_m) <= 256; } From ee7191a63094a79b097ba3dc163270ee3cfd0147 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Wed, 22 Jul 2026 12:01:46 -0700 Subject: [PATCH 03/11] fix build --- cmake/onnxruntime_cuda_source_filters.cmake | 19 +++++++++++++++++-- .../cuda/math/matmul_block_scaled_fp4.cc | 14 ++++++-------- 2 files changed, 23 insertions(+), 10 deletions(-) diff --git a/cmake/onnxruntime_cuda_source_filters.cmake b/cmake/onnxruntime_cuda_source_filters.cmake index dc9a6b3ce433a..5056a67a7f5a8 100644 --- a/cmake/onnxruntime_cuda_source_filters.cmake +++ b/cmake/onnxruntime_cuda_source_filters.cmake @@ -66,6 +66,19 @@ function(onnxruntime_extract_sm_specific_cuda_sources CU_SRC_LIST) set(_list "${${CU_SRC_LIST}}") + # MatMulBlockScaledFp4 native SM120 path must never be compiled in the default + # CUDA source list (e.g., SM86-only builds). Keep it only in SM120-specific + # object libraries when SM120 is requested. + set(_matmul_block_scaled_fp4_sm120_srcs) + foreach(_src IN LISTS _list) + if(_src MATCHES "matmul_block_scaled_fp4_sm120\\.cu$") + list(APPEND _matmul_block_scaled_fp4_sm120_srcs "${_src}") + endif() + endforeach() + if(_matmul_block_scaled_fp4_sm120_srcs) + list(REMOVE_ITEM _list ${_matmul_block_scaled_fp4_sm120_srcs}) + endif() + # Extract SM90 TMA WS generated files set(_sm90_srcs) if(ORT_HAS_SM90_OR_LATER) @@ -82,9 +95,11 @@ function(onnxruntime_extract_sm_specific_cuda_sources CU_SRC_LIST) # Extract SM120 TMA WS generated files set(_sm120_srcs) if("120" IN_LIST CMAKE_CUDA_ARCHITECTURES_ORIG) + if(_matmul_block_scaled_fp4_sm120_srcs) + list(APPEND _sm120_srcs ${_matmul_block_scaled_fp4_sm120_srcs}) + endif() foreach(_src IN LISTS _list) - if(_src MATCHES "moe_gemm_tma_ws_sm120_.*\\.generated\\.cu$" OR - _src MATCHES "matmul_block_scaled_fp4_sm120\\.cu$") + if(_src MATCHES "moe_gemm_tma_ws_sm120_.*\\.generated\\.cu$") list(APPEND _sm120_srcs "${_src}") endif() endforeach() diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index 38da1b064715b..7a20bc9e93449 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -176,23 +176,21 @@ Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { const int64_t rounded_m = RoundUp(m_i, 128); const int64_t rounded_n = RoundUp(n_i, 128); - auto a_packed = GetScratchBuffer(SafeInt(m_i) * SafeInt(k_i / 2), - context->GetComputeStream()); - auto a_scale = GetScratchBuffer(SafeInt(rounded_m) * SafeInt(k_scale_blocks), - context->GetComputeStream()); + onnxruntime::Stream* stream = GetComputeStream(context); + auto a_packed = GetScratchBuffer(SafeInt(m_i) * SafeInt(k_i / 2), stream); + auto a_scale = GetScratchBuffer(SafeInt(rounded_m) * SafeInt(k_scale_blocks), stream); IAllocatorUniquePtr b_scale; const void* b_scale_data = b_scale_prepacked_.get(); if (b_scale_data == nullptr) { - b_scale = GetScratchBuffer(SafeInt(rounded_n) * SafeInt(k_scale_blocks), - context->GetComputeStream()); + b_scale = GetScratchBuffer(SafeInt(rounded_n) * SafeInt(k_scale_blocks), stream); ORT_RETURN_IF_ERROR(LaunchRepackWeightScaleNvFp4ForNativeSm120( b_scale.get(), weight_scale->DataRaw(), n_i, k_i, SafeInt(block_size_), Stream(context))); b_scale_data = b_scale.get(); } - auto alpha = GetScratchBuffer(1, context->GetComputeStream()); + auto alpha = GetScratchBuffer(1, stream); const size_t workspace_size = GetMatMulBlockScaledFp4NativeSm120WorkspaceSize( m_i, n_i, k_i, std::is_same::value); - auto workspace = GetScratchBuffer(workspace_size, context->GetComputeStream()); + auto workspace = GetScratchBuffer(workspace_size, stream); ORT_RETURN_IF_ERROR(LaunchMatMulBlockScaledFp4NativeSm120( Y->MutableDataRaw(), From 8c3a8503bd5eed72dcdc6c81790fad986ed04f82 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Thu, 23 Jul 2026 13:54:20 -0700 Subject: [PATCH 04/11] MatMulBlockQuantizedFp4Weight --- cmake/onnxruntime_cuda_source_filters.cmake | 2 +- .../cuda/matmul_block_scaled_fp4.md | 12 +++++----- .../matmul_block_scaled_fp4_experiments.md | 6 ++--- .../contrib_ops/cuda/cuda_contrib_kernels.cc | 4 ++-- .../cuda/math/matmul_block_scaled_fp4.cc | 20 ++++++++-------- .../cuda/math/matmul_block_scaled_fp4.cu | 18 +++++++------- .../cuda/math/matmul_block_scaled_fp4.h | 10 ++++---- .../math/matmul_block_scaled_fp4_sm120.cu | 4 ++-- .../core/graph/contrib_ops/contrib_defs.cc | 2 +- onnxruntime/core/graph/contrib_ops/ms_opset.h | 4 ++-- .../matmul_block_scaled_fp4_test.cc | 24 +++++++++---------- .../profile_matmul_block_scaled.py | 6 ++--- 12 files changed, 56 insertions(+), 56 deletions(-) diff --git a/cmake/onnxruntime_cuda_source_filters.cmake b/cmake/onnxruntime_cuda_source_filters.cmake index 5056a67a7f5a8..66dadd5a02279 100644 --- a/cmake/onnxruntime_cuda_source_filters.cmake +++ b/cmake/onnxruntime_cuda_source_filters.cmake @@ -66,7 +66,7 @@ function(onnxruntime_extract_sm_specific_cuda_sources CU_SRC_LIST) set(_list "${${CU_SRC_LIST}}") - # MatMulBlockScaledFp4 native SM120 path must never be compiled in the default + # MatMulBlockQuantizedFp4Weight native SM120 path must never be compiled in the default # CUDA source list (e.g., SM86-only builds). Keep it only in SM120-specific # object libraries when SM120 is requested. set(_matmul_block_scaled_fp4_sm120_srcs) diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md index 07413932d266e..e16f39412e126 100644 --- a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md @@ -1,11 +1,11 @@ -# MatMulBlockScaledFp4 - CUDA Operator Documentation +# MatMulBlockQuantizedFp4Weight - CUDA Operator Documentation This document describes the CUDA execution-provider implementation of -**MatMulBlockScaledFp4** (`com.microsoft::MatMulBlockScaledFp4`): its tensor +**MatMulBlockQuantizedFp4Weight** (`com.microsoft::MatMulBlockQuantizedFp4Weight`): its tensor format, dispatch chain, native Blackwell path, prepacking behavior, and test / benchmark workflow. -MatMulBlockScaledFp4 computes `Y = A * dequant(B)^T (+ bias)` where `A` is +MatMulBlockQuantizedFp4Weight computes `Y = A * dequant(B)^T (+ bias)` where `A` is FP16 or BF16 and `B` is an `N x K` weight matrix stored as packed NVIDIA FP4 E2M1 values with block-wise E4M3 scales. The default semantics are weight-only FP4: activations stay FP16/BF16. An opt-in SM120 path quantizes @@ -79,7 +79,7 @@ native SM120 paths additionally require `block_size == 16` and `K % 32 == 0`. ## 3. Dispatch Chain -`MatMulBlockScaledFp4::ComputeImpl` tries the cheapest applicable path first: +`MatMulBlockQuantizedFp4Weight::ComputeImpl` tries the cheapest applicable path first: ```mermaid flowchart TD @@ -100,7 +100,7 @@ quantization, CUTLASS setup, and underutilized tensor-core GEMM work. ## 4. Decode Path - Fused GEMV -`LaunchMatMulBlockScaledFp4Gemv` is used when: +`LaunchMatMulBlockQuantizedFp4WeightGemv` is used when: - `0 < M <= 8`, - `block_size == 16`, @@ -213,7 +213,7 @@ Focused C++ tests: ```bash CUDA_VISIBLE_DEVICES=0 "$ORT_BUILD/onnxruntime_provider_test" \ - --gtest_filter='MatMulBlockScaledFp4OpTest.*' + --gtest_filter='MatMulBlockQuantizedFp4WeightOpTest.*' ``` Python harness examples: diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md index f70867e1b40b1..c0e10bff3e160 100644 --- a/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md @@ -1,7 +1,7 @@ -# MatMulBlockScaledFp4 - CUDA Experiments +# MatMulBlockQuantizedFp4Weight - CUDA Experiments This document records CUDA experiments for -**MatMulBlockScaledFp4** (`com.microsoft::MatMulBlockScaledFp4`) that are useful +**MatMulBlockQuantizedFp4Weight** (`com.microsoft::MatMulBlockQuantizedFp4Weight`) that are useful for future performance work but are not part of the final dispatch chain. Related documentation: @@ -173,7 +173,7 @@ Focused C++ tests: ```bash CUDA_VISIBLE_DEVICES=0 "$ORT_BUILD/onnxruntime_provider_test" \ - --gtest_filter='MatMulBlockScaledFp4OpTest.*' + --gtest_filter='MatMulBlockQuantizedFp4WeightOpTest.*' ``` --- diff --git a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc index 58807c7197844..b78c453bf92d7 100644 --- a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc @@ -180,7 +180,7 @@ class CUDA_ONNX_OP_TYPED_CLASS_NAME(1, float_float_MLFloat16, SimplifiedLayerNor class CUDA_ONNX_OP_TYPED_CLASS_NAME(1, MLFloat16_float_float, SimplifiedLayerNormalization); class CUDA_ONNX_OP_TYPED_CLASS_NAME(1, BFloat16_float_BFloat16, SimplifiedLayerNormalization); class CUDA_MS_OP_CLASS_NAME(1, Inverse); -class CUDA_MS_OP_CLASS_NAME(1, MatMulBlockScaledFp4); +class CUDA_MS_OP_CLASS_NAME(1, MatMulBlockQuantizedFp4Weight); class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, MatMulNBits); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, MatMulNBits); class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, MatMulNBits); @@ -442,7 +442,7 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index 7a20bc9e93449..c9a89bdc0e176 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -16,7 +16,7 @@ namespace onnxruntime::contrib::cuda { using namespace onnxruntime::cuda; ONNX_OPERATOR_KERNEL_EX( - MatMulBlockScaledFp4, + MatMulBlockQuantizedFp4Weight, kMSDomain, 1, kCudaExecutionProvider, @@ -25,7 +25,7 @@ ONNX_OPERATOR_KERNEL_EX( .TypeConstraint("T1", BuildKernelDefConstraints()) .TypeConstraint("T2", BuildKernelDefConstraints()) .TypeConstraint("T3", BuildKernelDefConstraints()), - MatMulBlockScaledFp4); + MatMulBlockQuantizedFp4Weight); namespace { @@ -41,7 +41,7 @@ int64_t RoundUp(int64_t value, int64_t alignment) { } // namespace -MatMulBlockScaledFp4::MatMulBlockScaledFp4(const OpKernelInfo& info) : CudaKernel(info) { +MatMulBlockQuantizedFp4Weight::MatMulBlockQuantizedFp4Weight(const OpKernelInfo& info) : CudaKernel(info) { ORT_ENFORCE(info.GetAttr("K", &K_).IsOK()); ORT_ENFORCE(info.GetAttr("N", &N_).IsOK()); block_size_ = info.GetAttrOrDefault("block_size", static_cast(16)); @@ -52,7 +52,7 @@ MatMulBlockScaledFp4::MatMulBlockScaledFp4(const OpKernelInfo& info) : CudaKerne sm_ = GetDeviceProp().major * 10 + GetDeviceProp().minor; } -Status MatMulBlockScaledFp4::PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, +Status MatMulBlockQuantizedFp4Weight::PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, bool& is_packed, PrePackedWeights* /*prepacked_weights*/) { is_packed = false; @@ -97,7 +97,7 @@ Status MatMulBlockScaledFp4::PrePack(const Tensor& tensor, int input_idx, Alloca } template -Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { +Status MatMulBlockQuantizedFp4Weight::ComputeImpl(OpKernelContext* context) const { typedef typename ToCudaType::MappedType CudaT; const Tensor* a = context->Input(0); @@ -153,7 +153,7 @@ Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { // [N, K] dequant scratch buffer and the cuBLAS GEMM (which is underutilized at M == 1). constexpr int kGemvMaxM = 8; if (m_i > 0 && m_i <= kGemvMaxM && block_size_ == 16 && (k_i % 32 == 0)) { - return LaunchMatMulBlockScaledFp4Gemv( + return LaunchMatMulBlockQuantizedFp4WeightGemv( Y->MutableDataRaw(), a->DataRaw(), b->DataRaw(), @@ -188,11 +188,11 @@ Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { b_scale_data = b_scale.get(); } auto alpha = GetScratchBuffer(1, stream); - const size_t workspace_size = GetMatMulBlockScaledFp4NativeSm120WorkspaceSize( + const size_t workspace_size = GetMatMulBlockQuantizedFp4WeightNativeSm120WorkspaceSize( m_i, n_i, k_i, std::is_same::value); auto workspace = GetScratchBuffer(workspace_size, stream); - ORT_RETURN_IF_ERROR(LaunchMatMulBlockScaledFp4NativeSm120( + ORT_RETURN_IF_ERROR(LaunchMatMulBlockQuantizedFp4WeightNativeSm120( Y->MutableDataRaw(), a->DataRaw(), b->DataRaw(), @@ -274,7 +274,7 @@ Status MatMulBlockScaledFp4::ComputeImpl(OpKernelContext* context) const { return Status::OK(); } -Status MatMulBlockScaledFp4::ComputeInternal(OpKernelContext* context) const { +Status MatMulBlockQuantizedFp4Weight::ComputeInternal(OpKernelContext* context) const { const Tensor* a = context->Input(0); if (a->IsDataType()) { return ComputeImpl(context); @@ -283,7 +283,7 @@ Status MatMulBlockScaledFp4::ComputeInternal(OpKernelContext* context) const { return ComputeImpl(context); } return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "MatMulBlockScaledFp4 only supports FP16 or BF16 activations."); + "MatMulBlockQuantizedFp4Weight only supports FP16 or BF16 activations."); } } // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu index a544a6d943530..7a2e61e85cace 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu @@ -140,7 +140,7 @@ __device__ __forceinline__ void LoadFp4Gemv32A(const nv_bfloat16* p } template -__global__ void MatMulBlockScaledFp4GemvKernel(T* __restrict__ y, +__global__ void MatMulBlockQuantizedFp4WeightGemvKernel(T* __restrict__ y, const T* __restrict__ a, const uint8_t* __restrict__ b_packed, const uint8_t* __restrict__ weight_scale, @@ -260,7 +260,7 @@ Status LaunchDequantizeNvFp4(void* b_dequant, ORT_UNUSED_PARAMETER(block_size); ORT_UNUSED_PARAMETER(is_bf16); ORT_UNUSED_PARAMETER(stream); - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockScaledFp4 requires CUDA 12.8 or newer for NVFP4 support."); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockQuantizedFp4Weight requires CUDA 12.8 or newer for NVFP4 support."); #endif } @@ -292,11 +292,11 @@ Status LaunchAddBiasNvFp4(void* y, ORT_UNUSED_PARAMETER(n); ORT_UNUSED_PARAMETER(is_bf16); ORT_UNUSED_PARAMETER(stream); - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockScaledFp4 requires CUDA 12.8 or newer for NVFP4 support."); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockQuantizedFp4Weight requires CUDA 12.8 or newer for NVFP4 support."); #endif } -Status LaunchMatMulBlockScaledFp4Gemv(void* y, +Status LaunchMatMulBlockQuantizedFp4WeightGemv(void* y, const void* a, const void* b_packed, const void* weight_scale, @@ -315,8 +315,8 @@ Status LaunchMatMulBlockScaledFp4Gemv(void* y, // This kernel is hard-coded for block_size == 16 and assumes K is a multiple of 32 so that each // warp lane always owns a full 32-element slice (and one E4M3 scale per 16-element block). Guard // against misuse if this helper is ever reused outside the callers that already check these. - ORT_RETURN_IF_NOT(block_size == 16, "MatMulBlockScaledFp4 GEMV requires block_size == 16, got ", block_size, "."); - ORT_RETURN_IF_NOT(k % 32 == 0, "MatMulBlockScaledFp4 GEMV requires K divisible by 32, got ", k, "."); + ORT_RETURN_IF_NOT(block_size == 16, "MatMulBlockQuantizedFp4Weight GEMV requires block_size == 16, got ", block_size, "."); + ORT_RETURN_IF_NOT(k % 32 == 0, "MatMulBlockQuantizedFp4Weight GEMV requires K divisible by 32, got ", k, "."); const int k_blocks = (k + block_size - 1) / block_size; constexpr int kWarpsPerBlock = 8; const dim3 threads{32, kWarpsPerBlock}; @@ -325,11 +325,11 @@ Status LaunchMatMulBlockScaledFp4Gemv(void* y, const uint8_t* bp = reinterpret_cast(b_packed); const uint8_t* ws = reinterpret_cast(weight_scale); if (is_bf16) { - MatMulBlockScaledFp4GemvKernel<<>>( + MatMulBlockQuantizedFp4WeightGemvKernel<<>>( reinterpret_cast(y), reinterpret_cast(a), bp, ws, weight_scale_2, reinterpret_cast(bias), m, n, k, k_blocks); } else { - MatMulBlockScaledFp4GemvKernel<<>>( + MatMulBlockQuantizedFp4WeightGemvKernel<<>>( reinterpret_cast(y), reinterpret_cast(a), bp, ws, weight_scale_2, reinterpret_cast(bias), m, n, k, k_blocks); } @@ -347,7 +347,7 @@ Status LaunchMatMulBlockScaledFp4Gemv(void* y, ORT_UNUSED_PARAMETER(block_size); ORT_UNUSED_PARAMETER(is_bf16); ORT_UNUSED_PARAMETER(stream); - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockScaledFp4 requires CUDA 12.8 or newer for NVFP4 support."); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "MatMulBlockQuantizedFp4Weight requires CUDA 12.8 or newer for NVFP4 support."); #endif } diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h index 3da8d52ac8e6c..149400ad6005c 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h @@ -14,9 +14,9 @@ namespace onnxruntime::contrib::cuda { // The weight is dequantized to the activation type (FP16/BF16) and multiplied with the FP16/BF16 // activation via cuBLAS. This path works on any CUDA architecture (including Hopper/SM90) because // it does not rely on native NVFP4 block-scaled tensor cores (SM100/SM120 only). -class MatMulBlockScaledFp4 final : public onnxruntime::cuda::CudaKernel { +class MatMulBlockQuantizedFp4Weight final : public onnxruntime::cuda::CudaKernel { public: - explicit MatMulBlockScaledFp4(const OpKernelInfo& info); + explicit MatMulBlockQuantizedFp4Weight(const OpKernelInfo& info); Status PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, bool& is_packed, PrePackedWeights* prepacked_weights) override; @@ -62,7 +62,7 @@ Status LaunchAddBiasNvFp4(void* y, // uint8 (raw E4M3 bytes), weight_scale_2 is a device fp32 scalar, bias is an optional [N] vector // (may be null). Output y is [M, N] in the activation type. Requires block_size == 16 and // k % 32 == 0. Runs on any architecture with NVFP4 conversion intrinsics (CUDA >= 12.8). -Status LaunchMatMulBlockScaledFp4Gemv(void* y, +Status LaunchMatMulBlockQuantizedFp4WeightGemv(void* y, const void* a, const void* b_packed, const void* weight_scale, @@ -86,7 +86,7 @@ Status LaunchRepackWeightScaleNvFp4ForNativeSm120(void* b_scale, int block_size, cudaStream_t stream); -Status LaunchMatMulBlockScaledFp4NativeSm120(void* y, +Status LaunchMatMulBlockQuantizedFp4WeightNativeSm120(void* y, const void* a, const void* b_packed, const void* weight_scale, @@ -104,6 +104,6 @@ Status LaunchMatMulBlockScaledFp4NativeSm120(void* y, void* workspace, size_t workspace_size, cudaStream_t stream); -size_t GetMatMulBlockScaledFp4NativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16); +size_t GetMatMulBlockQuantizedFp4WeightNativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16); } // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu index 59911bfe3f0e1..4f4945955da21 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu @@ -287,7 +287,7 @@ Status DispatchRunGemm(void* y, } // namespace -size_t GetMatMulBlockScaledFp4NativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16) { +size_t GetMatMulBlockQuantizedFp4WeightNativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16) { return is_bf16 ? DispatchWorkspaceSize(m, n, k) : DispatchWorkspaceSize(m, n, k); } @@ -314,7 +314,7 @@ Status LaunchRepackWeightScaleNvFp4ForNativeSm120(void* b_scale, return CUDA_CALL(cudaGetLastError()); } -Status LaunchMatMulBlockScaledFp4NativeSm120(void* y, +Status LaunchMatMulBlockQuantizedFp4WeightNativeSm120(void* y, const void* a, const void* b_packed, const void* weight_scale, diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index 8013429fd7604..f9fd961ed636a 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -3006,7 +3006,7 @@ ONNX_MS_OPERATOR_SET_SCHEMA(GemmFloat8, 1, })); ONNX_MS_OPERATOR_SET_SCHEMA( - MatMulBlockScaledFp4, 1, + MatMulBlockQuantizedFp4Weight, 1, OpSchema() .SetDoc(R"DOC(Weight-only NVFP4 (E2M1) matrix multiplication. diff --git a/onnxruntime/core/graph/contrib_ops/ms_opset.h b/onnxruntime/core/graph/contrib_ops/ms_opset.h index fbc259575708f..206dccc43d491 100644 --- a/onnxruntime/core/graph/contrib_ops/ms_opset.h +++ b/onnxruntime/core/graph/contrib_ops/ms_opset.h @@ -121,7 +121,7 @@ class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GemmFastGelu); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, DecoderMaskedSelfAttention); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, DecoderMaskedMultiHeadAttention); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GemmFloat8); -class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulBlockScaledFp4); +class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulBlockQuantizedFp4Weight); class OpSet_Microsoft_ver1 { public: @@ -237,7 +237,7 @@ class OpSet_Microsoft_ver1 { fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); - fn(GetOpSchema()); + fn(GetOpSchema()); } }; } // namespace contrib diff --git a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc index 7e1c0c03172d7..7ec03968eeb47 100644 --- a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc @@ -21,9 +21,9 @@ namespace onnxruntime::test { // A -> [M, K] all ones per row scaled by (m + 1); weights are constant per row, so the // operator must reproduce Y[m, n] = W_val[n] * sum_k A[m, k]. -TEST(MatMulBlockScaledFp4OpTest, WeightOnlyBasicFp16) { +TEST(MatMulBlockQuantizedFp4WeightOpTest, WeightOnlyBasicFp16) { if (!HasCudaEnvironment(800)) { - GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; } constexpr int64_t m = 2; @@ -49,7 +49,7 @@ TEST(MatMulBlockScaledFp4OpTest, WeightOnlyBasicFp16) { // W[0, :] = 1.0, W[1, :] = 2.0; sum_k A[0, :] = 16, sum_k A[1, :] = 32. std::vector expected = {16.0f, 32.0f, 32.0f, 64.0f}; - OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); test.AddAttribute("K", k); test.AddAttribute("N", n); test.AddAttribute("block_size", 16); @@ -66,9 +66,9 @@ TEST(MatMulBlockScaledFp4OpTest, WeightOnlyBasicFp16) { // Exercises non-unit per-block E4M3 scales, a global weight_scale_2, negative weights, bias and // a skipped optional input_scale, with BF16 activations/output. -TEST(MatMulBlockScaledFp4OpTest, WeightOnlyScalesBiasBf16) { +TEST(MatMulBlockQuantizedFp4WeightOpTest, WeightOnlyScalesBiasBf16) { if (!HasCudaEnvironment(800)) { - GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; } constexpr int64_t m = 1; @@ -91,7 +91,7 @@ TEST(MatMulBlockScaledFp4OpTest, WeightOnlyScalesBiasBf16) { std::vector bias = {1.0f, 2.0f}; std::vector expected = {97.0f, -46.0f}; - OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); test.AddAttribute("K", k); test.AddAttribute("N", n); test.AddAttribute("block_size", 16); @@ -111,9 +111,9 @@ TEST(MatMulBlockScaledFp4OpTest, WeightOnlyScalesBiasBf16) { // Exercises the fused decode GEMV fast path (small M) with a multi-block K (K = 64 == 4 blocks, // K % 32 == 0), FP16 activations. Weights are constant per row so Y[m, n] = W_val[n] * sum_k A[m, k]. -TEST(MatMulBlockScaledFp4OpTest, GemvDecodeMultiBlockFp16) { +TEST(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeMultiBlockFp16) { if (!HasCudaEnvironment(800)) { - GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; } constexpr int64_t m = 2; @@ -140,7 +140,7 @@ TEST(MatMulBlockScaledFp4OpTest, GemvDecodeMultiBlockFp16) { // W[0, :] = 1.0, W[1, :] = 2.0; sum_k A[0, :] = 64, sum_k A[1, :] = 128. std::vector expected = {64.0f, 128.0f, 128.0f, 256.0f}; - OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); test.AddAttribute("K", k); test.AddAttribute("N", n); test.AddAttribute("block_size", 16); @@ -158,9 +158,9 @@ TEST(MatMulBlockScaledFp4OpTest, GemvDecodeMultiBlockFp16) { // Exercises the fused decode GEMV fast path with M == 1, per-block scales, a global weight_scale_2, // negative weights and bias (BF16). K = 32 == 2 blocks, K % 32 == 0 -> GEMV path. -TEST(MatMulBlockScaledFp4OpTest, GemvDecodeScalesBiasBf16) { +TEST(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeScalesBiasBf16) { if (!HasCudaEnvironment(800)) { - GTEST_SKIP() << "CUDA device is required for MatMulBlockScaledFp4."; + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; } constexpr int64_t m = 1; @@ -184,7 +184,7 @@ TEST(MatMulBlockScaledFp4OpTest, GemvDecodeScalesBiasBf16) { std::vector bias = {1.0f, 2.0f}; std::vector expected = {193.0f, -94.0f}; - OpTester test("MatMulBlockScaledFp4", 1, onnxruntime::kMSDomain); + OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); test.AddAttribute("K", k); test.AddAttribute("N", n); test.AddAttribute("block_size", 16); diff --git a/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py b/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py index b2c1644b924b0..4745648272f7b 100644 --- a/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py +++ b/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py @@ -4,7 +4,7 @@ # -------------------------------------------------------------------------- """ -Accuracy and latency harness for the CUDA MatMulBlockScaledFp4 contrib op. +Accuracy and latency harness for the CUDA MatMulBlockQuantizedFp4Weight contrib op. The script builds a single-node com.microsoft contrib-op model, binds CUDA tensors with I/O binding, compares the output with an FP32 dequantized reference, and prints one JSON @@ -184,7 +184,7 @@ def _make_fp4_model(case: Case, b_packed: torch.Tensor, weight_scale: torch.Tens initializers.append(_make_float_initializer("bias", bias, activation_onnx_type)) node = helper.make_node( - "MatMulBlockScaledFp4", + "MatMulBlockQuantizedFp4Weight", inputs, ["Y"], domain="com.microsoft", @@ -194,7 +194,7 @@ def _make_fp4_model(case: Case, b_packed: torch.Tensor, weight_scale: torch.Tens ) graph_inputs = [helper.make_tensor_value_info("A", activation_onnx_type, [case.m, case.k])] graph_outputs = [helper.make_tensor_value_info("Y", activation_onnx_type, [case.m, case.n])] - return _model_bytes([node], graph_inputs, graph_outputs, initializers, "MatMulBlockScaledFp4_Profile") + return _model_bytes([node], graph_inputs, graph_outputs, initializers, "MatMulBlockQuantizedFp4Weight_Profile") def _fp4_reference(a: torch.Tensor, b_dequantized: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor: From 48146642918e28aa16c367ecab929e8f54e778be Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Thu, 23 Jul 2026 17:00:20 -0700 Subject: [PATCH 05/11] lintrunner --- .../cuda/math/matmul_block_scaled_fp4.cc | 2 +- .../cuda/math/matmul_block_scaled_fp4.cu | 40 ++++++------- .../cuda/math/matmul_block_scaled_fp4.h | 56 +++++++++---------- .../math/matmul_block_scaled_fp4_sm120.cu | 34 +++++------ 4 files changed, 66 insertions(+), 66 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index c9a89bdc0e176..291221944427c 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -53,7 +53,7 @@ MatMulBlockQuantizedFp4Weight::MatMulBlockQuantizedFp4Weight(const OpKernelInfo& } Status MatMulBlockQuantizedFp4Weight::PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, - bool& is_packed, PrePackedWeights* /*prepacked_weights*/) { + bool& is_packed, PrePackedWeights* /*prepacked_weights*/) { is_packed = false; #if defined(ORT_ENABLE_BLOCKQUANT_SM120) diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu index 7a2e61e85cace..9a1b9ebaaef73 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu @@ -141,15 +141,15 @@ __device__ __forceinline__ void LoadFp4Gemv32A(const nv_bfloat16* p template __global__ void MatMulBlockQuantizedFp4WeightGemvKernel(T* __restrict__ y, - const T* __restrict__ a, - const uint8_t* __restrict__ b_packed, - const uint8_t* __restrict__ weight_scale, - const float* __restrict__ weight_scale_2, - const T* __restrict__ bias, - int m, - int n, - int k, - int k_blocks) { + const T* __restrict__ a, + const uint8_t* __restrict__ b_packed, + const uint8_t* __restrict__ weight_scale, + const float* __restrict__ weight_scale_2, + const T* __restrict__ bias, + int m, + int n, + int k, + int k_blocks) { const int lane = threadIdx.x; // 0..31 const int col = blockIdx.x * blockDim.y + threadIdx.y; // n const int row = blockIdx.y; // m @@ -297,17 +297,17 @@ Status LaunchAddBiasNvFp4(void* y, } Status LaunchMatMulBlockQuantizedFp4WeightGemv(void* y, - const void* a, - const void* b_packed, - const void* weight_scale, - const float* weight_scale_2, - const void* bias, - int m, - int n, - int k, - int block_size, - bool is_bf16, - cudaStream_t stream) { + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const void* bias, + int m, + int n, + int k, + int block_size, + bool is_bf16, + cudaStream_t stream) { #if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 if (m <= 0 || n <= 0 || k <= 0) { return Status::OK(); diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h index 149400ad6005c..b48a947991873 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h @@ -63,17 +63,17 @@ Status LaunchAddBiasNvFp4(void* y, // (may be null). Output y is [M, N] in the activation type. Requires block_size == 16 and // k % 32 == 0. Runs on any architecture with NVFP4 conversion intrinsics (CUDA >= 12.8). Status LaunchMatMulBlockQuantizedFp4WeightGemv(void* y, - const void* a, - const void* b_packed, - const void* weight_scale, - const float* weight_scale_2, - const void* bias, - int m, - int n, - int k, - int block_size, - bool is_bf16, - cudaStream_t stream); + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const void* bias, + int m, + int n, + int k, + int block_size, + bool is_bf16, + cudaStream_t stream); // Native Blackwell SM120 NVFP4 x NVFP4 GEMM path. The caller provides scratch buffers for // packed activation FP4, swizzled A/B scale tensors, alpha, and CUTLASS workspace. A is [M, K] @@ -87,23 +87,23 @@ Status LaunchRepackWeightScaleNvFp4ForNativeSm120(void* b_scale, cudaStream_t stream); Status LaunchMatMulBlockQuantizedFp4WeightNativeSm120(void* y, - const void* a, - const void* b_packed, - const void* weight_scale, - const float* weight_scale_2, - const float* input_scale, - void* a_packed, - void* a_scale, - const void* b_scale, - float* alpha, - int m, - int n, - int k, - int block_size, - bool is_bf16, - void* workspace, - size_t workspace_size, - cudaStream_t stream); + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const float* input_scale, + void* a_packed, + void* a_scale, + const void* b_scale, + float* alpha, + int m, + int n, + int k, + int block_size, + bool is_bf16, + void* workspace, + size_t workspace_size, + cudaStream_t stream); size_t GetMatMulBlockQuantizedFp4WeightNativeSm120WorkspaceSize(int m, int n, int k, bool is_bf16); } // namespace onnxruntime::contrib::cuda diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu index 4f4945955da21..e9a0e5513fe07 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu @@ -315,23 +315,23 @@ Status LaunchRepackWeightScaleNvFp4ForNativeSm120(void* b_scale, } Status LaunchMatMulBlockQuantizedFp4WeightNativeSm120(void* y, - const void* a, - const void* b_packed, - const void* weight_scale, - const float* weight_scale_2, - const float* input_scale, - void* a_packed, - void* a_scale, - const void* b_scale, - float* alpha, - int m, - int n, - int k, - int block_size, - bool is_bf16, - void* workspace, - size_t workspace_size, - cudaStream_t stream) { + const void* a, + const void* b_packed, + const void* weight_scale, + const float* weight_scale_2, + const float* input_scale, + void* a_packed, + void* a_scale, + const void* b_scale, + float* alpha, + int m, + int n, + int k, + int block_size, + bool is_bf16, + void* workspace, + size_t workspace_size, + cudaStream_t stream) { ORT_UNUSED_PARAMETER(workspace_size); ORT_UNUSED_PARAMETER(weight_scale); ORT_RETURN_IF_NOT(block_size == kScaleVectorSize, From fb7ae6b9227aae646fa7543e8541320a2d45cf3b Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Fri, 24 Jul 2026 06:25:36 +0000 Subject: [PATCH 06/11] fix plugin ep build --- onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index 291221944427c..c011034671cd4 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -228,7 +228,7 @@ Status MatMulBlockQuantizedFp4Weight::ComputeImpl(OpKernelContext* context) cons // Dequantize the packed NVFP4 weight into a scratch [N, K] buffer of the activation type. IAllocatorUniquePtr b_dequant = GetScratchBuffer(SafeInt(N_) * SafeInt(K_), - context->GetComputeStream()); + GetComputeStream(context)); ORT_RETURN_IF_ERROR(LaunchDequantizeNvFp4( b_dequant.get(), b->DataRaw(), From ec3ab4316cd7c3bddce003e5c8371d7bba6de1a8 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Fri, 24 Jul 2026 06:31:30 +0000 Subject: [PATCH 07/11] docs: update QMoE NVFP4 reference --- docs/ContribOperators.md | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index f7b9aa0a01b77..16d90e886b59f 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -4943,7 +4943,7 @@ This version of the operator has been available since version 1 of the 'com.micr
normalize_routing_weights : int
Whether to normalize routing weights
quant_type : string
-
Quantization type: 'int' for integer quantization (default), 'fp4' for MXFP4 quantization, 'fp8' for FP8 e4m3 weight-only quantization, or 'wfp4afp8' for MXFP4 weight with FP8 activation. When quant_type is 'fp4', weights are stored in MXFP4 format (2 values per byte), fc*_scales inputs contain MXFP4 block scales, and fc*_global_scale inputs must be provided.
+
Quantization type: 'int' for integer quantization (default), 'fp4' for MXFP4 quantization, 'nvfp4' for NVFP4 quantization, 'fp8' for FP8 e4m3 weight-only quantization, or 'wfp4afp8' for MXFP4 weight with FP8 activation. When quant_type is 'fp4' or 'nvfp4', weights are stored in E2M1 FP4 format (2 values per byte), fc*_scales inputs contain the FP4 block scales, and fc*_global_scale inputs must be provided. 'fp4' uses Float8E8M0 block scales with block_size 32; 'nvfp4' uses Float8E4M3FN block scales with block_size 16.
swiglu_fusion : int
0: not fused, 1: fused and interleaved. 2: fused and not interleaved.
swiglu_limit : float
@@ -4964,13 +4964,13 @@ This version of the operator has been available since version 1 of the 'com.micr
fc1_experts_weights : T1
3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / pack_size), The fusion_size is 2 for fused swiglu, or 1 otherwise. The pack_size is 8 / expert_weight_bits.
fc1_scales (optional) : T2
-
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, fusion_size * inter_size), or a 3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, fusion_size * inter_size, hidden_size / 32). Not used for quant_type='fp8'.
+
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, fusion_size * inter_size), or a 3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, fusion_size * inter_size, hidden_size / 32). For quant_type='nvfp4', this is a float8e4m3fn NVFP4 block-scale tensor with shape (num_experts, fusion_size * inter_size, hidden_size / 16). Not used for quant_type='fp8'.
fc1_experts_bias (optional) : T
2D optional tensor with shape (num_experts, fusion_size * inter_size)
fc2_experts_weights : T1
3D tensor with shape (num_experts, hidden_size, inter_size / pack_size)
fc2_scales (optional) : T2
-
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, hidden_size), or a 3D tensor with shape (num_experts, hidden_size, inter_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, hidden_size, inter_size / 32). Not used for quant_type='fp8'.
+
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, hidden_size), or a 3D tensor with shape (num_experts, hidden_size, inter_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, hidden_size, inter_size / 32). For quant_type='nvfp4', this is a float8e4m3fn NVFP4 block-scale tensor with shape (num_experts, hidden_size, inter_size / 16). Not used for quant_type='fp8'.
fc2_experts_bias (optional) : T
2D optional tensor with shape (num_experts, hidden_size)
fc3_experts_weights (optional) : T1
@@ -4988,9 +4988,9 @@ This version of the operator has been available since version 1 of the 'com.micr
router_weights (optional) : T
2D optional tensor with shape (num_tokens, num_experts). When provided, router_probs is used only for Top-K expert selection, and router_weights is used for aggregating expert outputs (the values at the selected expert indices are gathered and used as mixing weights). This enables DeepSeek-style noaux_tc routing where different tensors are used for selection and aggregation. When not provided, router_probs is used for both selection and aggregation (backward compatible).
fc1_global_scale (optional) : T4
-
1D optional tensor with shape (num_experts,). Per-expert global weight scale for FC1. Required when quant_type is 'fp4', 'fp8', or 'wfp4afp8'.
+
1D optional tensor with shape (num_experts,). Per-expert global weight scale for FC1. Required when quant_type is 'fp4', 'nvfp4', 'fp8', or 'wfp4afp8'.
fc2_global_scale (optional) : T4
-
1D optional tensor with shape (num_experts,). Per-expert global weight scale for FC2. Required when quant_type is 'fp4', 'fp8', or 'wfp4afp8'.
+
1D optional tensor with shape (num_experts,). Per-expert global weight scale for FC2. Required when quant_type is 'fp4', 'nvfp4', 'fp8', or 'wfp4afp8'.
fc1_act_scale (optional) : T4
1D optional tensor with shape (1,) or (num_experts,). Activation scale for FC1 FP8 activation modes.
fc2_act_scale (optional) : T4
@@ -5015,8 +5015,8 @@ This version of the operator has been available since version 1 of the 'com.micr
Constrain input and output types to float tensors.
T1 : tensor(uint8), tensor(float8e4m3fn)
Constrain quantized weight types. Integer and FP4 weights use uint8. FP8 weights use float8e4m3fn.
-
T2 : tensor(float), tensor(float16), tensor(bfloat16), tensor(float8e8m0)
-
Constrain scale types. Float tensors are used for integer quantization scales. Float8e8m0 tensors are used for MXFP block scales.
+
T2 : tensor(float), tensor(float16), tensor(bfloat16), tensor(float8e8m0), tensor(float8e4m3fn)
+
Constrain scale types. Float tensors are used for integer quantization scales. Float8e8m0 tensors are used for MXFP4 block scales; float8e4m3fn tensors are used for NVFP4 block scales.
T4 : tensor(float)
Constrain FP4 global scale type to float32 tensors.
From c8c2e117c4a495eb31436f1b50a45dca50220497 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Fri, 24 Jul 2026 02:15:59 -0700 Subject: [PATCH 08/11] Update generated operator docs --- docs/ContribOperators.md | 65 ++++++++++++++++++++++++++++++++++++++++ docs/OperatorKernels.md | 1 + 2 files changed, 66 insertions(+) diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index 16d90e886b59f..ce31181aec046 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -52,6 +52,7 @@ Do not modify directly.* * com.microsoft.Irfft * com.microsoft.LinearAttention * com.microsoft.LongformerAttention + * com.microsoft.MatMulBlockQuantizedFp4Weight * com.microsoft.MatMulBnb4 * com.microsoft.MatMulFpQ4 * com.microsoft.MatMulInteger16 @@ -2909,6 +2910,70 @@ This version of the operator has been available since version 1 of the 'com.micr +### **com.microsoft.MatMulBlockQuantizedFp4Weight** + + Weight-only NVFP4 (E2M1) matrix multiplication. + + The weight tensor B is stored as packed NVFP4: two E2M1 values per byte (low nibble first). + The dequantized weight value is `e2m1(B) * weight_scale_2 * e4m3(weight_scale[n, k / block_size])`, + where `weight_scale` holds one E4M3 scale per `block_size` (default 16) consecutive K values and + `weight_scale_2` is a single global fp32 scale. The weight is dequantized to the activation type + (FP16/BF16) and multiplied with the FP16/BF16 activation. This path is architecture independent and + runs on Hopper (SM90) as well as Blackwell. + +#### Version + +This version of the operator has been available since version 1 of the 'com.microsoft' operator set. + +#### Attributes + +
+
K : int (required)
+
Inner (contraction) dimension: the number of logical columns of the unpacked weight.
+
N : int (required)
+
Number of output columns, i.e. the number of rows of the packed weight.
+
block_size : int
+
Number of consecutive K values that share one E4M3 weight scale. Default 16.
+
+ +#### Inputs (4 - 6) + +
+
A : T
+
Row-major FP16/BF16 activation of shape [..., K].
+
B : T1
+
Packed NVFP4 weight of shape [N, K/2] stored as uint8 (two E2M1 values per byte, low nibble first).
+
weight_scale : T2
+
Per-block E4M3 weight scales of shape [N, ceil(K / block_size)] stored as raw uint8 bytes.
+
weight_scale_2 : T3
+
Global fp32 weight scale (scalar).
+
input_scale (optional) : T3
+
Optional global fp32 activation scale (scalar). Accepted for parity with quantized checkpoints; it is a no-op on the weight-only FP16/BF16 path and is reserved for the native NVFP4 path on Blackwell.
+
bias (optional) : T
+
Optional bias of shape [N].
+
+ +#### Outputs + +
+
Y : T
+
Output of shape [..., N] in the activation type.
+
+ +#### Type Constraints + +
+
T : tensor(float16), tensor(bfloat16)
+
Constrain activation, bias and output to FP16 or BF16.
+
T1 : tensor(uint8)
+
Constrain packed NVFP4 weight to uint8.
+
T2 : tensor(uint8)
+
Constrain E4M3 weight scales to uint8.
+
T3 : tensor(float)
+
Constrain scalar scales to FP32.
+
+ + ### **com.microsoft.MatMulBnb4** MatMulBnb4 is a MatMul with weight quantized with 4 bits using either FP4 or NF4 data type (https://arxiv.org/pdf/2305.14314.pdf). It does Matrix Multiplication like MatMul (https://github.com/onnx/onnx/blob/main/docs/Operators.md#matmul) with differences: diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md index eda8bc19a1ae0..356075089e578 100644 --- a/docs/OperatorKernels.md +++ b/docs/OperatorKernels.md @@ -1099,6 +1099,7 @@ The **OpSet Version** column uses the following notation: |Irfft|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)| |LinearAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* past_state:**S**
*in* decay:**T**
*in* beta:**T**
*out* output:**T**
*out* present_state:**S**|1+|**T** = tensor(float), tensor(float16)| |LongformerAttention|*in* input:**T**
*in* weight:**T**
*in* bias:**T**
*in* mask:**T**
*in* global_weight:**T**
*in* global_bias:**T**
*in* global:**G**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| +|MatMulBlockQuantizedFp4Weight|*in* A:**T**
*in* B:**T1**
*in* weight_scale:**T2**
*in* weight_scale_2:**T3**
*in* input_scale:**T3**
*in* bias:**T**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(float16)
**T1** = tensor(uint8)
**T2** = tensor(uint8)
**T3** = tensor(float)| |MatMulBnb4|*in* A:**T1**
*in* B:**T2**
*in* absmax:**T1**
*out* Y:**T1**|1+|**T1** = tensor(bfloat16), tensor(float), tensor(float16)
**T2** = tensor(uint8)| |MatMulNBits|*in* A:**T1**
*in* B:**T2**
*in* scales:**T1**
*in* zero_points:**T3**
*in* g_idx:**T4**
*in* bias:**T1**
*out* Y:**T1**|1+|**T1** = tensor(bfloat16), tensor(float), tensor(float16)
**T2** = tensor(uint8)
**T3** = tensor(bfloat16), tensor(float), tensor(float16), tensor(uint8)| |MoE|*in* input:**T**
*in* router_probs:**T**
*in* fc1_experts_weights:**T**
*in* fc1_experts_bias:**T**
*in* fc2_experts_weights:**T**
*in* fc2_experts_bias:**T**
*in* fc3_experts_weights:**T**
*in* fc3_experts_bias:**T**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| From 600efc399b7b0244011c82e35b7445f2618dd938 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Sun, 26 Jul 2026 05:10:07 +0000 Subject: [PATCH 09/11] remove N and K attributes --- docs/ContribOperators.md | 7 +-- .../cuda/matmul_block_scaled_fp4.md | 10 +++- .../cuda/math/matmul_block_scaled_fp4.cc | 59 ++++++++++--------- .../cuda/math/matmul_block_scaled_fp4.h | 2 - .../core/graph/contrib_ops/contrib_defs.cc | 29 ++++----- .../matmul_block_scaled_fp4_test.cc | 8 --- .../profile_matmul_block_scaled.py | 2 - 7 files changed, 57 insertions(+), 60 deletions(-) diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index 05d1cff2880aa..95f0698534a74 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -2921,6 +2921,9 @@ This version of the operator has been available since version 1 of the 'com.micr `weight_scale_2` is a single global fp32 scale. The weight is dequantized to the activation type (FP16/BF16) and multiplied with the FP16/BF16 activation. This path is architecture independent and runs on Hopper (SM90) as well as Blackwell. + + The output columns `N` and the contraction dimension `K` are derived from the weight shape: + `N = B.shape[0]` and `K = 2 * B.shape[1]`. `K` must therefore be even. #### Version @@ -2929,10 +2932,6 @@ This version of the operator has been available since version 1 of the 'com.micr #### Attributes
-
K : int (required)
-
Inner (contraction) dimension: the number of logical columns of the unpacked weight.
-
N : int (required)
-
Number of output columns, i.e. the number of rows of the packed weight.
block_size : int
Number of consecutive K values that share one E4M3 weight scale. Default 16.
diff --git a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md index e16f39412e126..cb3ccde088f5d 100644 --- a/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md @@ -39,10 +39,11 @@ Source files: | Attribute | Meaning | |-----------|---------| -| `K` | Input feature dimension: columns of `A` and logical columns of `B`. | -| `N` | Output feature dimension: rows of logical `B`. | | `block_size` | Quantization group size along `K`. Current CUDA paths are optimized for `16`; default is `16`. | +`N` and `K` are not attributes. They are derived from the weight shape: +`N = B.shape[0]` and `K = 2 * B.shape[1]`. + | Input | Index | Type | Notes | |-------|-------|------|-------| | `A` | 0 | FP16 or BF16 | Activation tensor with last dimension `K`. Leading dimensions are flattened into `M`. | @@ -175,7 +176,10 @@ select this path. `PrePack` handles input index `2` (`weight_scale`) only for the eligible native SM120 path. It converts the original `[N, K / 16]` E4M3 scale tensor into the -SM120 swizzled scale layout once and stores it in `b_scale_prepacked_`. +SM120 swizzled scale layout once and stores it in `b_scale_prepacked_`. Because +the weight tensor is not visible in `PrePack`, `N` and `K` are recovered from the +scale shape itself (`N = scale.shape[0]`, `K = scale.shape[1] * block_size`), +which is exact for the `K % block_size == 0` shapes this path requires. `is_packed` deliberately remains `false`: the original `weight_scale` input must stay available because the decode GEMV and default dequant+cuBLAS paths still diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index c011034671cd4..68567f45e1731 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -42,13 +42,8 @@ int64_t RoundUp(int64_t value, int64_t alignment) { } // namespace MatMulBlockQuantizedFp4Weight::MatMulBlockQuantizedFp4Weight(const OpKernelInfo& info) : CudaKernel(info) { - ORT_ENFORCE(info.GetAttr("K", &K_).IsOK()); - ORT_ENFORCE(info.GetAttr("N", &N_).IsOK()); block_size_ = info.GetAttrOrDefault("block_size", static_cast(16)); - ORT_ENFORCE(K_ > 0, "K must be positive, got ", K_); - ORT_ENFORCE(N_ > 0, "N must be positive, got ", N_); ORT_ENFORCE(block_size_ > 0, "block_size must be positive, got ", block_size_); - ORT_ENFORCE(K_ % 2 == 0, "K must be even for packed NVFP4 weights, got ", K_); sm_ = GetDeviceProp().major * 10 + GetDeviceProp().minor; } @@ -58,18 +53,25 @@ Status MatMulBlockQuantizedFp4Weight::PrePack(const Tensor& tensor, int input_id #if defined(ORT_ENABLE_BLOCKQUANT_SM120) if (input_idx != kWeightScaleInputIndex || !IsNativeSm120Fp4Enabled() || sm_ < 120 || sm_ >= 130 || - block_size_ != 16 || K_ % 32 != 0 || N_ % 32 != 0) { + block_size_ != 16) { return Status::OK(); } - const int64_t k_blocks = K_ / 16; + // N and K are derived from the scale shape [N, K / block_size]. This is exact only when K is a + // multiple of block_size, which the native SM120 path requires anyway (K % 32 == 0). const auto& scale_shape = tensor.Shape(); - ORT_RETURN_IF_NOT(scale_shape.NumDimensions() == 2 && scale_shape[0] == N_ && scale_shape[1] == k_blocks, - "weight_scale must have shape [N, K/16] = [", N_, ", ", k_blocks, "], got ", - scale_shape.ToString(), "."); + if (scale_shape.NumDimensions() != 2) { + return Status::OK(); + } + const int64_t n = scale_shape[0]; + const int64_t k_blocks = scale_shape[1]; + const int64_t k = k_blocks * block_size_; + if (k % 32 != 0 || n % 32 != 0) { + return Status::OK(); + } const int64_t rounded_k_blocks = RoundUp(k_blocks, 4); - const int64_t rounded_n = RoundUp(N_, 128); + const int64_t rounded_n = RoundUp(n, 128); b_scale_prepacked_ = IAllocator::MakeUniquePtr( alloc, SafeInt(rounded_n) * SafeInt(rounded_k_blocks), true); @@ -77,7 +79,7 @@ Status MatMulBlockQuantizedFp4Weight::PrePack(const Tensor& tensor, int input_id const void* weight_scale = tensor.DataRaw(); IAllocatorUniquePtr weight_scale_device; if (tensor.Location().device.Type() != OrtDevice::GPU) { - const size_t weight_scale_bytes = SafeInt(N_) * SafeInt(k_blocks); + const size_t weight_scale_bytes = SafeInt(n) * SafeInt(k_blocks); weight_scale_device = IAllocator::MakeUniquePtr(alloc, weight_scale_bytes, true); CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(weight_scale_device.get(), weight_scale, weight_scale_bytes, cudaMemcpyDefault, stream)); @@ -85,7 +87,7 @@ Status MatMulBlockQuantizedFp4Weight::PrePack(const Tensor& tensor, int input_id } ORT_RETURN_IF_ERROR(LaunchRepackWeightScaleNvFp4ForNativeSm120( - b_scale_prepacked_.get(), weight_scale, SafeInt(N_), SafeInt(K_), SafeInt(block_size_), stream)); + b_scale_prepacked_.get(), weight_scale, SafeInt(n), SafeInt(k), SafeInt(block_size_), stream)); CUDA_RETURN_IF_ERROR(cudaStreamSynchronize(stream)); #else ORT_UNUSED_PARAMETER(tensor); @@ -109,18 +111,21 @@ Status MatMulBlockQuantizedFp4Weight::ComputeImpl(OpKernelContext* context) cons const auto& a_shape = a->Shape(); ORT_ENFORCE(a_shape.NumDimensions() >= 1, "A must have rank at least 1."); - ORT_ENFORCE(a_shape[a_shape.NumDimensions() - 1] == K_, - "A's last dimension (", a_shape[a_shape.NumDimensions() - 1], ") must equal K (", K_, ")."); - const int64_t k_packed = K_ / 2; - const int64_t k_blocks = (K_ + block_size_ - 1) / block_size_; + // N and K are derived from the packed weight shape [N, K/2] instead of from attributes. const auto& b_shape = b->Shape(); - ORT_ENFORCE(b_shape.NumDimensions() == 2 && b_shape[0] == N_ && b_shape[1] == k_packed, - "B must have shape [N, K/2] = [", N_, ", ", k_packed, "], got ", b_shape.ToString(), "."); + ORT_ENFORCE(b_shape.NumDimensions() == 2, "B must have shape [N, K/2], got ", b_shape.ToString(), "."); + const int64_t n = b_shape[0]; + const int64_t k = b_shape[1] * 2; + ORT_ENFORCE(a_shape[a_shape.NumDimensions() - 1] == k, + "A's last dimension (", a_shape[a_shape.NumDimensions() - 1], + ") must equal K (", k, ") derived from B's shape ", b_shape.ToString(), "."); + + const int64_t k_blocks = (k + block_size_ - 1) / block_size_; const auto& weight_scale_shape = weight_scale->Shape(); - ORT_ENFORCE(weight_scale_shape.NumDimensions() == 2 && weight_scale_shape[0] == N_ && + ORT_ENFORCE(weight_scale_shape.NumDimensions() == 2 && weight_scale_shape[0] == n && weight_scale_shape[1] == k_blocks, - "weight_scale must have shape [N, ceil(K/block_size)] = [", N_, ", ", k_blocks, "], got ", + "weight_scale must have shape [N, ceil(K/block_size)] = [", n, ", ", k_blocks, "], got ", weight_scale_shape.ToString(), "."); ORT_ENFORCE(weight_scale_2->Shape().Size() == 1, "weight_scale_2 must be a scalar."); if (input_scale != nullptr) { @@ -129,14 +134,14 @@ Status MatMulBlockQuantizedFp4Weight::ComputeImpl(OpKernelContext* context) cons // weight-only FP16/BF16 activation path keeps full-precision activations. } if (bias != nullptr) { - ORT_ENFORCE(bias->Shape().NumDimensions() == 1 && bias->Shape()[0] == N_, - "bias must have shape [N] = [", N_, "], got ", bias->Shape().ToString(), "."); + ORT_ENFORCE(bias->Shape().NumDimensions() == 1 && bias->Shape()[0] == n, + "bias must have shape [N] = [", n, "], got ", bias->Shape().ToString(), "."); } constexpr bool transa = false; constexpr bool transb = true; MatMulComputeHelper helper; - TensorShape b_logical_shape({N_, K_}); + TensorShape b_logical_shape({n, k}); ORT_RETURN_IF_ERROR(helper.Compute(a_shape, b_logical_shape, transa, transb)); Tensor* Y = context->Output(0, helper.OutputShape()); @@ -227,15 +232,15 @@ Status MatMulBlockQuantizedFp4Weight::ComputeImpl(OpKernelContext* context) cons #endif // Dequantize the packed NVFP4 weight into a scratch [N, K] buffer of the activation type. - IAllocatorUniquePtr b_dequant = GetScratchBuffer(SafeInt(N_) * SafeInt(K_), + IAllocatorUniquePtr b_dequant = GetScratchBuffer(SafeInt(n) * SafeInt(k), GetComputeStream(context)); ORT_RETURN_IF_ERROR(LaunchDequantizeNvFp4( b_dequant.get(), b->DataRaw(), weight_scale->DataRaw(), weight_scale_2->Data(), - SafeInt(N_), - SafeInt(K_), + SafeInt(n), + SafeInt(k), SafeInt(block_size_), std::is_same::value, Stream(context))); diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h index b48a947991873..f62fe7bf147a5 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h @@ -27,8 +27,6 @@ class MatMulBlockQuantizedFp4Weight final : public onnxruntime::cuda::CudaKernel template Status ComputeImpl(OpKernelContext* context) const; - int64_t K_; - int64_t N_; int64_t block_size_; int sm_{0}; IAllocatorUniquePtr b_scale_prepacked_; diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index 47a9630b662ac..3ad978455dd5f 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -3015,11 +3015,10 @@ The dequantized weight value is `e2m1(B) * weight_scale_2 * e4m3(weight_scale[n, where `weight_scale` holds one E4M3 scale per `block_size` (default 16) consecutive K values and `weight_scale_2` is a single global fp32 scale. The weight is dequantized to the activation type (FP16/BF16) and multiplied with the FP16/BF16 activation. This path is architecture independent and -runs on Hopper (SM90) as well as Blackwell.)DOC") - .Attr("K", "Inner (contraction) dimension: the number of logical columns of the unpacked weight.", - AttributeProto::INT) - .Attr("N", "Number of output columns, i.e. the number of rows of the packed weight.", - AttributeProto::INT) +runs on Hopper (SM90) as well as Blackwell. + +The output columns `N` and the contraction dimension `K` are derived from the weight shape: +`N = B.shape[0]` and `K = 2 * B.shape[1]`. `K` must therefore be even.)DOC") .Attr("block_size", "Number of consecutive K values that share one E4M3 weight scale. Default 16.", AttributeProto::INT, static_cast(16)) .Input(0, "A", "Row-major FP16/BF16 activation of shape [..., K].", "T") @@ -3042,23 +3041,25 @@ runs on Hopper (SM90) as well as Blackwell.)DOC") .TypeConstraint("T3", {"tensor(float)"}, "Constrain scalar scales to FP32.") .TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) { propagateElemTypeFromInputToOutput(ctx, 0, 0); - if (!hasInputShape(ctx, 0)) { + if (!hasNInputShapes(ctx, 2)) { return; } const auto& a_shape = getInputShape(ctx, 0); - if (a_shape.dim_size() < 1) { - fail_shape_inference("A must have rank at least 1."); + const auto& b_shape = getInputShape(ctx, 1); + if (a_shape.dim_size() < 1 || b_shape.dim_size() != 2) { + fail_shape_inference("A must have rank at least 1 and B must have rank 2."); + } + // B is packed two E2M1 values per byte, so the logical K is twice B's last dimension. + const auto& a_k = a_shape.dim(a_shape.dim_size() - 1); + if (a_k.has_dim_value() && b_shape.dim(1).has_dim_value() && + a_k.dim_value() != 2 * b_shape.dim(1).dim_value()) { + fail_shape_inference("A and B have incompatible K dimensions."); } ONNX_NAMESPACE::TensorShapeProto output_shape; for (int i = 0; i < a_shape.dim_size() - 1; ++i) { *output_shape.add_dim() = a_shape.dim(i); } - const auto* n_attr = ctx.getAttribute("N"); - if (n_attr != nullptr && n_attr->has_i()) { - output_shape.add_dim()->set_dim_value(n_attr->i()); - } else { - output_shape.add_dim(); - } + *output_shape.add_dim() = b_shape.dim(0); updateOutputShape(ctx, 0, output_shape); })); diff --git a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc index 7ec03968eeb47..8fcd3f49b9fec 100644 --- a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc @@ -50,8 +50,6 @@ TEST(MatMulBlockQuantizedFp4WeightOpTest, WeightOnlyBasicFp16) { std::vector expected = {16.0f, 32.0f, 32.0f, 64.0f}; OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); - test.AddAttribute("K", k); - test.AddAttribute("N", n); test.AddAttribute("block_size", 16); test.AddInput("A", {m, k}, FloatsToMLFloat16s(a)); test.AddInput("B", {n, k / 2}, b); @@ -92,8 +90,6 @@ TEST(MatMulBlockQuantizedFp4WeightOpTest, WeightOnlyScalesBiasBf16) { std::vector expected = {97.0f, -46.0f}; OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); - test.AddAttribute("K", k); - test.AddAttribute("N", n); test.AddAttribute("block_size", 16); test.AddInput("A", {m, k}, FloatsToBFloat16s(a)); test.AddInput("B", {n, k / 2}, b); @@ -141,8 +137,6 @@ TEST(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeMultiBlockFp16) { std::vector expected = {64.0f, 128.0f, 128.0f, 256.0f}; OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); - test.AddAttribute("K", k); - test.AddAttribute("N", n); test.AddAttribute("block_size", 16); test.AddInput("A", {m, k}, FloatsToMLFloat16s(a)); test.AddInput("B", {n, k / 2}, b); @@ -185,8 +179,6 @@ TEST(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeScalesBiasBf16) { std::vector expected = {193.0f, -94.0f}; OpTester test("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); - test.AddAttribute("K", k); - test.AddAttribute("N", n); test.AddAttribute("block_size", 16); test.AddInput("A", {m, k}, FloatsToBFloat16s(a)); test.AddInput("B", {n, k / 2}, b); diff --git a/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py b/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py index f6e33aae5ede4..8ac3c126476ee 100644 --- a/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py +++ b/onnxruntime/test/python/contrib_ops/profile_matmul_block_scaled.py @@ -193,8 +193,6 @@ def _make_fp4_model(case: Case, b_packed: torch.Tensor, weight_scale: torch.Tens inputs, ["Y"], domain="com.microsoft", - K=case.k, - N=case.n, block_size=block_size, ) graph_inputs = [helper.make_tensor_value_info("A", activation_onnx_type, [case.m, case.k])] From e24d029804d1259fe97e8898164725cf13ddd089 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Sun, 26 Jul 2026 08:23:08 +0000 Subject: [PATCH 10/11] fix(cuda): handle K == 0 and gate FP4 tests on CUDA 12.8+ Address review feedback on MatMulBlockQuantizedFp4Weight: - K == 0 satisfied the `k % 32 == 0` condition used to select the GEMV and native SM120 fast paths. The GEMV launcher returned OK without writing Y (uninitialized output) and the SM120 path would have launched grids with grid.x == 0. Handle the degenerate empty-reduction case explicitly: zero Y and add the optional bias, so the fast paths can assume K > 0. - The test file was only guarded by USE_CUDA, but the operator requires CUDA 12.8+ for the NVFP4 intrinsics. Gate on CUDA_VERSION >= 12080 and include so the macro is actually visible in the test TU. --- .../cuda/math/matmul_block_scaled_fp4.cc | 16 ++++++++++++++++ .../contrib_ops/matmul_block_scaled_fp4_test.cc | 10 ++++++++-- 2 files changed, 24 insertions(+), 2 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc index 68567f45e1731..8663febcd9700 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -153,6 +153,22 @@ Status MatMulBlockQuantizedFp4Weight::ComputeImpl(OpKernelContext* context) cons const int n_i = SafeInt(helper.N()); const int k_i = SafeInt(helper.K()); + // Degenerate K == 0 reduces to an empty sum, so Y is zero plus the optional bias. Handle it + // here so that the specialized paths below can assume K > 0 (K % 32 == 0 is also true for 0). + if (k_i == 0) { + CUDA_RETURN_IF_ERROR(cudaMemsetAsync(Y->MutableDataRaw(), 0, Y->SizeInBytes(), Stream(context))); + if (bias != nullptr) { + ORT_RETURN_IF_ERROR(LaunchAddBiasNvFp4( + Y->MutableDataRaw(), + bias->DataRaw(), + m_i, + n_i, + std::is_same::value, + Stream(context))); + } + return Status::OK(); + } + // Decode fast path: for small M (autoregressive generation) this is a memory-bound GEMV. // A fused warp-per-column kernel reads the packed NVFP4 weight directly, avoiding both the // [N, K] dequant scratch buffer and the cuBLAS GEMM (which is underutilized at M == 1). diff --git a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc index 8fcd3f49b9fec..243887d5495a5 100644 --- a/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc +++ b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc @@ -1,6 +1,12 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +#if defined(USE_CUDA) +// Needed for the CUDA_VERSION check below. MatMulBlockQuantizedFp4Weight relies on the NVFP4 +// conversion intrinsics that are only available in CUDA 12.8 and newer. +#include +#endif + #include "gtest/gtest.h" #include "test/common/cuda_op_test_utils.h" #include "test/common/tensor_op_test_utils.h" @@ -9,7 +15,7 @@ namespace onnxruntime::test { -#if defined(USE_CUDA) +#if defined(USE_CUDA) && defined(CUDA_VERSION) && CUDA_VERSION >= 12080 // NVFP4 (E2M1) 4-bit magnitude nibble encodings (sign bit is 0x8): // +0.0 -> 0x0, +0.5 -> 0x1, +1.0 -> 0x2, +1.5 -> 0x3, @@ -194,6 +200,6 @@ TEST(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeScalesBiasBf16) { test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); } -#endif // USE_CUDA +#endif // USE_CUDA && defined(CUDA_VERSION) && CUDA_VERSION >= 12080 } // namespace onnxruntime::test From eb36211a69a9ccf26142f3c127263c2b87658656 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Mon, 27 Jul 2026 08:27:57 +0000 Subject: [PATCH 11/11] replace __nv_cvt_fp4x2_to_halfraw2 by bit trick --- .../cuda/math/matmul_block_scaled_fp4.cu | 101 +++++++++++------- 1 file changed, 60 insertions(+), 41 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu index 9a1b9ebaaef73..2066bef6b209f 100644 --- a/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu @@ -6,6 +6,7 @@ #include #include #include +#include #if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 #include @@ -108,36 +109,53 @@ __global__ void AddBiasKernel(T* __restrict__ y, const T* __restrict__ bias, int // per-block E4M3 scales are folded in per half. The global fp32 scale is applied // once after the warp reduction. Runs on any architecture with NVFP4 conversion // intrinsics (CUDA >= 12.8), including SM90 and SM120. +// ----------------------------------------------------------------------------- +// Fast NVFP4 (E2M1) -> half / bfloat16 conversion. +// +// __nv_cvt_fp4x2_to_halfraw2() is emulated in software on pre-Blackwell parts and +// costs ~10 ALU ops per pair, which dominates the decode GEMV below (measured 3.5x +// slowdown on H200). Instead build the target float directly from the code bits: +// for c = s e1 e0 m the pattern (s << 15) | ((c & 7) << ) has +// exponent field == e and mantissa == m/2, i.e. exactly value * 2^-(bias - 1). +// Multiplying by 2^(bias - 1) afterwards recovers the value, and the e == 0 +// subnormal encodings fall out correctly as well. Verified exhaustively against the +// intrinsic for all 256 packed byte values, for both half and bfloat16. template -__device__ __forceinline__ void LoadFp4Gemv32A(const T* ptr, float (&out)[32]); +struct Fp4Cvt; template <> -__device__ __forceinline__ void LoadFp4Gemv32A(const half* ptr, float (&out)[32]) { - const uint4* p = reinterpret_cast(ptr); -#pragma unroll - for (int j = 0; j < 4; ++j) { - const uint4 raw = p[j]; - const half* v = reinterpret_cast(&raw); -#pragma unroll - for (int i = 0; i < 8; ++i) { - out[j * 8 + i] = __half2float(v[i]); - } +struct Fp4Cvt { + using T2 = half2; + static __device__ __forceinline__ T2 Raw(uint32_t b) { + const uint32_t lo = ((b & 0x07u) << 9) | ((b & 0x08u) << 12); + const uint32_t hi = ((b & 0x70u) << 5) | ((b & 0x80u) << 8); + const uint32_t bits = lo | (hi << 16); + T2 r; + memcpy(&r, &bits, sizeof(r)); + return r; } -} + static __device__ __forceinline__ T2 Scale() { return __float2half2_rn(16384.0f); } // 2^14 + static __device__ __forceinline__ T2 Mul(T2 a, T2 b) { return __hmul2(a, b); } + static __device__ __forceinline__ float2 ToFloat2(T2 v) { return __half22float2(v); } +}; template <> -__device__ __forceinline__ void LoadFp4Gemv32A(const nv_bfloat16* ptr, float (&out)[32]) { - const uint4* p = reinterpret_cast(ptr); -#pragma unroll - for (int j = 0; j < 4; ++j) { - const uint4 raw = p[j]; - const nv_bfloat16* v = reinterpret_cast(&raw); -#pragma unroll - for (int i = 0; i < 8; ++i) { - out[j * 8 + i] = __bfloat162float(v[i]); - } +struct Fp4Cvt { + using T2 = nv_bfloat162; + static __device__ __forceinline__ T2 Raw(uint32_t b) { + const uint32_t lo = ((b & 0x07u) << 6) | ((b & 0x08u) << 12); + const uint32_t hi = ((b & 0x70u) << 2) | ((b & 0x80u) << 8); + const uint32_t bits = lo | (hi << 16); + T2 r; + memcpy(&r, &bits, sizeof(r)); + return r; } -} + static __device__ __forceinline__ T2 Scale() { + return __float2bfloat162_rn(85070591730234615865843651857942052864.0f); // 2^126 + } + static __device__ __forceinline__ T2 Mul(T2 a, T2 b) { return __hmul2(a, b); } + static __device__ __forceinline__ float2 ToFloat2(T2 v) { return __bfloat1622float2(v); } +}; template __global__ void MatMulBlockQuantizedFp4WeightGemvKernel(T* __restrict__ y, @@ -150,6 +168,9 @@ __global__ void MatMulBlockQuantizedFp4WeightGemvKernel(T* __restrict__ y, int n, int k, int k_blocks) { + using Cvt = Fp4Cvt; + using T2 = typename Cvt::T2; + const int lane = threadIdx.x; // 0..31 const int col = blockIdx.x * blockDim.y + threadIdx.y; // n const int row = blockIdx.y; // m @@ -161,6 +182,7 @@ __global__ void MatMulBlockQuantizedFp4WeightGemvKernel(T* __restrict__ y, const uint8_t* b_row = b_packed + static_cast(col) * (k >> 1); const uint8_t* ws_row = weight_scale + static_cast(col) * k_blocks; + const T2 up = Cvt::Scale(); constexpr int kBlockSize = 16; constexpr int kElemsPerLane = 32; // two 16-element blocks const int stride = 32 * kElemsPerLane; // 1024 elements per warp iteration @@ -171,31 +193,28 @@ __global__ void MatMulBlockQuantizedFp4WeightGemvKernel(T* __restrict__ y, if (koff < k) { const uint4 packed = *reinterpret_cast(b_row + (koff >> 1)); const uint8_t* bytes = reinterpret_cast(&packed); - float b_vals[32]; -#pragma unroll - for (int i = 0; i < 16; ++i) { - const __half2_raw hr = __nv_cvt_fp4x2_to_halfraw2( - static_cast<__nv_fp4x2_storage_t>(bytes[i]), __NV_E2M1); - const float2 f = __half22float2(__half2(hr)); - b_vals[i * 2] = f.x; - b_vals[i * 2 + 1] = f.y; - } - float a_vals[32]; - LoadFp4Gemv32A(a_row + koff, a_vals); + const uint4* ap = reinterpret_cast(a_row + koff); + uint4 a0 = ap[0], a1 = ap[1], a2 = ap[2], a3 = ap[3]; + const T2* av[4] = {reinterpret_cast(&a0), reinterpret_cast(&a1), + reinterpret_cast(&a2), reinterpret_cast(&a3)}; - const int kb0 = koff / kBlockSize; - const int kb1 = kb0 + 1; + // fp32 accumulation per 16-element scale block, matching the reference path. float p0 = 0.0f; float p1 = 0.0f; #pragma unroll for (int i = 0; i < 16; ++i) { - p0 += a_vals[i] * b_vals[i]; - } -#pragma unroll - for (int i = 16; i < 32; ++i) { - p1 += a_vals[i] * b_vals[i]; + const T2 bb = Cvt::Mul(Cvt::Raw(bytes[i]), up); + const float2 pv = Cvt::ToFloat2(Cvt::Mul(av[i >> 2][i & 3], bb)); + if (i < 8) { + p0 += pv.x + pv.y; + } else { + p1 += pv.x + pv.y; + } } + + const int kb0 = koff / kBlockSize; + const int kb1 = kb0 + 1; const float s0 = __half2float(__nv_cvt_fp8_to_halfraw( static_cast<__nv_fp8_storage_t>(ws_row[kb0]), __NV_E4M3)); const float s1 = __half2float(__nv_cvt_fp8_to_halfraw(