diff --git a/cmake/onnxruntime_cuda_source_filters.cmake b/cmake/onnxruntime_cuda_source_filters.cmake index 5b15603aa0acd..99b9cdd4690bc 100644 --- a/cmake/onnxruntime_cuda_source_filters.cmake +++ b/cmake/onnxruntime_cuda_source_filters.cmake @@ -67,6 +67,19 @@ function(onnxruntime_extract_sm_specific_cuda_sources CU_SRC_LIST) set(_list "${${CU_SRC_LIST}}") + # 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) + 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) @@ -83,6 +96,9 @@ 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 OR "121" 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$") list(APPEND _sm120_srcs "${_src}") diff --git a/cmake/onnxruntime_providers_cuda.cmake b/cmake/onnxruntime_providers_cuda.cmake index 1b5c7ab5b2a6d..be26843c57596 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 53edf1ae184c7..bde686630ecd9 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/ContribOperators.md b/docs/ContribOperators.md index 7b3d56aa9a59b..95f0698534a74 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.MatMulBlockQuantizedFp8Weight * com.microsoft.MatMulBnb4 * com.microsoft.MatMulFpQ4 @@ -2910,6 +2911,68 @@ 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. + + 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 + +This version of the operator has been available since version 1 of the 'com.microsoft' operator set. + +#### Attributes + +
+
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.MatMulBlockQuantizedFp8Weight** Weight-only block-scaled FP8 (E4M3) matrix multiplication. diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md index d31eeec1101c5..5f2da22aabe65 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)| |MatMulBlockQuantizedFp8Weight|*in* A:**T**
*in* B:**T1**
*in* b_scale:**T2**
*in* a_scale:**T2**
*in* bias:**T**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(float16)
**T1** = tensor(float8e4m3fn)
**T2** = 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)| 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..cb3ccde088f5d --- /dev/null +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4.md @@ -0,0 +1,251 @@ +# MatMulBlockQuantizedFp4Weight - CUDA Operator Documentation + +This document describes the CUDA execution-provider implementation of +**MatMulBlockQuantizedFp4Weight** (`com.microsoft::MatMulBlockQuantizedFp4Weight`): its tensor +format, dispatch chain, native Blackwell path, prepacking behavior, and test / +benchmark workflow. + +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 +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 | +|-----------|---------| +| `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`. | +| `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 + +`MatMulBlockQuantizedFp4Weight::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 + +`LaunchMatMulBlockQuantizedFp4WeightGemv` 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_`. 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 +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 + +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 "$ORT_BUILD/onnxruntime_provider_test" \ + --gtest_filter='MatMulBlockQuantizedFp4WeightOpTest.*' +``` + +Python harness examples: + +```bash +# Decode GEMV +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="$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="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + 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 +``` + +After rebuilding `libonnxruntime_providers_cuda.so`, sync the provider into the +Python load locations before Python benchmarks: + +```bash +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 new file mode 100644 index 0000000000000..c0e10bff3e160 --- /dev/null +++ b/docs/contrib_ops/cuda/matmul_block_scaled_fp4_experiments.md @@ -0,0 +1,189 @@ +# MatMulBlockQuantizedFp4Weight - CUDA Experiments + +This document records CUDA experiments for +**MatMulBlockQuantizedFp4Weight** (`com.microsoft::MatMulBlockQuantizedFp4Weight`) 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 + +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 "$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="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + 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="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + 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="$ORT_BUILD" CUDA_VISIBLE_DEVICES=0 \ + ORT_MATMUL_BLOCK_SCALED_FP4_NATIVE_SM120=1 \ + 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 "$ORT_BUILD/onnxruntime_provider_test" \ + --gtest_filter='MatMulBlockQuantizedFp4WeightOpTest.*' +``` + +--- + +## 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 77d74dd3d9a24..1c898ea5912a8 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, MatMulBlockQuantizedFp4Weight); #if !defined(DISABLE_FLOAT8_TYPES) class CUDA_MS_OP_CLASS_NAME(1, MatMulBlockQuantizedFp8Weight); #endif @@ -444,6 +445,7 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, #if !defined(DISABLE_FLOAT8_TYPES) BuildKernelCreateInfo, #endif 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..8663febcd9700 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cc @@ -0,0 +1,310 @@ +// 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( + MatMulBlockQuantizedFp4Weight, + kMSDomain, + 1, + kCudaExecutionProvider, + (*KernelDefBuilder::Create()) + .TypeConstraint("T", BuildKernelDefConstraints()) + .TypeConstraint("T1", BuildKernelDefConstraints()) + .TypeConstraint("T2", BuildKernelDefConstraints()) + .TypeConstraint("T3", BuildKernelDefConstraints()), + MatMulBlockQuantizedFp4Weight); + +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 + +MatMulBlockQuantizedFp4Weight::MatMulBlockQuantizedFp4Weight(const OpKernelInfo& info) : CudaKernel(info) { + block_size_ = info.GetAttrOrDefault("block_size", static_cast(16)); + ORT_ENFORCE(block_size_ > 0, "block_size must be positive, got ", block_size_); + sm_ = GetDeviceProp().major * 10 + GetDeviceProp().minor; +} + +Status MatMulBlockQuantizedFp4Weight::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) { + return Status::OK(); + } + + // 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(); + 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); + 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 MatMulBlockQuantizedFp4Weight::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."); + + // 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 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 && + 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."); + // 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().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}); + 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()); + + // 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). + constexpr int kGemvMaxM = 8; + if (m_i > 0 && m_i <= kGemvMaxM && block_size_ == 16 && (k_i % 32 == 0)) { + return LaunchMatMulBlockQuantizedFp4WeightGemv( + 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); + + 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), 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, stream); + 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(LaunchMatMulBlockQuantizedFp4WeightNativeSm120( + 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), + GetComputeStream(context)); + 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 MatMulBlockQuantizedFp4Weight::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, + "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 new file mode 100644 index 0000000000000..2066bef6b209f --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.cu @@ -0,0 +1,373 @@ +// 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 +#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. +// ----------------------------------------------------------------------------- +// 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 +struct Fp4Cvt; + +template <> +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 <> +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, + 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) { + 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 + 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; + + 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 + + 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); + + 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)}; + + // 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) { + 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( + 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, "MatMulBlockQuantizedFp4Weight 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, "MatMulBlockQuantizedFp4Weight requires CUDA 12.8 or newer for NVFP4 support."); +#endif +} + +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) { +#if defined(CUDA_VERSION) && CUDA_VERSION >= 12080 + 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, "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}; + 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) { + MatMulBlockQuantizedFp4WeightGemvKernel<<>>( + reinterpret_cast(y), reinterpret_cast(a), bp, ws, weight_scale_2, + reinterpret_cast(bias), m, n, k, k_blocks); + } else { + MatMulBlockQuantizedFp4WeightGemvKernel<<>>( + 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, "MatMulBlockQuantizedFp4Weight 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..f62fe7bf147a5 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4.h @@ -0,0 +1,107 @@ +// 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 MatMulBlockQuantizedFp4Weight final : public onnxruntime::cuda::CudaKernel { + public: + explicit MatMulBlockQuantizedFp4Weight(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 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 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); + +// 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 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); +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 new file mode 100644 index 0000000000000..e9a0e5513fe07 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/math/matmul_block_scaled_fp4_sm120.cu @@ -0,0 +1,364 @@ +// 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) { + // 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; +} + +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 GetMatMulBlockQuantizedFp4WeightNativeSm120WorkspaceSize(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 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) { + 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 18f19eefc296a..3ad978455dd5f 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -3005,6 +3005,64 @@ 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( + MatMulBlockQuantizedFp4Weight, 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. + +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") + .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 (!hasNInputShapes(ctx, 2)) { + return; + } + const auto& a_shape = getInputShape(ctx, 0); + 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); + } + *output_shape.add_dim() = b_shape.dim(0); + updateOutputShape(ctx, 0, output_shape); + })); + ONNX_MS_OPERATOR_SET_SCHEMA( MatMulBlockQuantizedFp8Weight, 1, OpSchema() diff --git a/onnxruntime/core/graph/contrib_ops/ms_opset.h b/onnxruntime/core/graph/contrib_ops/ms_opset.h index b7634daafb84f..84ee0acd0bc8e 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, MatMulBlockQuantizedFp4Weight); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulBlockQuantizedFp8Weight); class OpSet_Microsoft_ver1 { @@ -237,6 +238,7 @@ class OpSet_Microsoft_ver1 { fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); + fn(GetOpSchema()); fn(GetOpSchema()); } }; 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..243887d5495a5 --- /dev/null +++ b/onnxruntime/test/contrib_ops/matmul_block_scaled_fp4_test.cc @@ -0,0 +1,205 @@ +// 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" +#include "test/providers/provider_test_utils.h" +#include "test/unittest_util/conversion.h" + +namespace onnxruntime::test { + +#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, +// +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(MatMulBlockQuantizedFp4WeightOpTest, WeightOnlyBasicFp16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; + } + + 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("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); + 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(MatMulBlockQuantizedFp4WeightOpTest, WeightOnlyScalesBiasBf16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; + } + + 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("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); + 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(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeMultiBlockFp16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; + } + + 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("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); + 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(MatMulBlockQuantizedFp4WeightOpTest, GemvDecodeScalesBiasBf16) { + if (!HasCudaEnvironment(800)) { + GTEST_SKIP() << "CUDA device is required for MatMulBlockQuantizedFp4Weight."; + } + + 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("MatMulBlockQuantizedFp4Weight", 1, onnxruntime::kMSDomain); + 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 && defined(CUDA_VERSION) && CUDA_VERSION >= 12080 + +} // 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 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])]