GEMM + sReLU (SM100)
GEMM + sReLU (SM100)
This is an experimental API and subject to change.
Overview
Block-scaled GEMM + sReLU fusion: A persistent, batched dense GEMM on NVIDIA Blackwell GPUs (SM100+) that supports block-scaled FP4 and FP8 inputs and produces both the full GEMM result C and a probability-gated squared-ReLU output D in a single kernel launch.
- Inputs: quantized
AandB, scale-factor tensorsSFAandSFB, and a per-row probability tensorprob - Outputs: full GEMM result
C, squared-ReLU outputD, and optional output scale factorsSFD/Amax
Shapes
-
Inputs
A: shape(M, K, L)B: shape(N, K, L)SFA: shape(32, 4, ceil_div(M, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), L)SFB: shape(32, 4, ceil_div(N, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), L)prob: shape(M, 1, L)
-
Outputs
C: shape(M, N, L)D: shape(M, N, L)SFD: shape(32, 4, ceil_div(M, 128), 4, ceil_div(ceil_div(N, sf_vec_size), 4), L)whenDis FP8Amax: shape(1,)when FP4 input is written to fp16/bf16/fp32 output
L is the batch dimension.
Equations
Let A_hat and B_hat denote the dequantized inputs from (A, SFA) and (B, SFB).
When D is FP8, the kernel also emits output scale factors SFD using the provided norm_const_tensor. When FP4 input is written to a higher-precision D, the kernel can also emit Amax.
Diagram
API Usage
The tensor parameters are type-erased: torch tensors and JAX arrays are both accepted (torch is only imported when torch tensors/dtypes are passed, jax only when JAX arrays are passed). Dtype parameters accept torch dtypes, numpy/ml_dtypes dtypes, dtype name strings, or cutlass types. The JAX contract matches gemm_amax (see gemm_amax.md “Using JAX arrays”): A/B k-major (M, K, 1)/(N, K, 1), outputs n-major only, batch L == 1, scale-factor tensors accepted in the physical C-contiguous atom shape (L, MN', K', 32, 4, 4); the eager entry points run on the CUDA legacy default stream (synchronize before reading outputs). For jitted JAX programs use the jax.jit-compatible XLA custom-call entry point (gemm_srelu_jax_sm100 / gemm_dsrelu_jax_sm100, built on cudnn.jax.call); the kernels’ optional parameters are compile-time constants inside its adapter.
High-level wrapper
Class API
Parameters
Input/Output tensors
- Input tensor A:
a_tensor(wrapper) orsample_a/a_tensor(class)- Shape:
(M, K, L) - Dtype:
{float4_e2m1fn_x2, uint8, float8_e4m3fn, float8_e5m2} - Note:
uint8is interpreted as packedfloat4_e2m1fn_x2(FP4x2) data, not integer quantization
- Shape:
- Input tensor B:
b_tensor(wrapper) orsample_b/b_tensor(class)- Shape:
(N, K, L) - Dtype: Must match
A
- Shape:
- Input tensor SFA:
sfa_tensor(wrapper) orsample_sfa/sfa_tensor(class)- Shape:
(32, 4, ceil_div(M, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), L) - Dtype:
{float8_e8m0fnu, float8_e4m3fn}
- Shape:
- Input tensor SFB:
sfb_tensor(wrapper) orsample_sfb/sfb_tensor(class)- Shape:
(32, 4, ceil_div(N, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), L) - Dtype: Must match
SFA
- Shape:
- Input tensor prob:
prob_tensor(wrapper) orsample_prob/prob_tensor(class)- Shape:
(M, 1, L) - Dtype:
float32
- Shape:
- Output tensor C:
result["c_tensor"](wrapper) orsample_c/c_tensor(class)- Shape:
(M, N, L) - Dtype:
{float16, bfloat16, float32, float8_e4m3fn, float8_e5m2}
- Shape:
- Output tensor D:
result["d_tensor"](wrapper) orsample_d/d_tensor(class)- Shape:
(M, N, L) - Dtype:
{float16, bfloat16, float32, float8_e4m3fn, float8_e5m2}
- Shape:
- Output tensor SFD:
result["sfd_tensor"](wrapper) orsample_sfd/sfd_tensor(class)- Shape:
(32, 4, ceil_div(M, 128), 4, ceil_div(ceil_div(N, sf_vec_size), 4), L) - Dtype: Must match
SFA - Required when
Dis FP8
- Shape:
- Output tensor Amax:
result["amax_tensor"](wrapper) orsample_amax/amax_tensor(class)- Shape:
(1,) - Dtype:
float32 - Allocated by the wrapper for FP4 input with fp16/bf16/fp32
D
- Shape:
- Input tensor Norm Const:
norm_const_tensor(wrapper) orsample_norm_const/norm_const_tensor(class)- Shape:
(1,) - Dtype:
float32 - Required when
Dis FP8
- Shape:
Common parameters
alpha: float- Scalar multiplier applied to the GEMM result before the sReLU epilogue. Default:
1.0
- Scalar multiplier applied to the GEMM result before the sReLU epilogue. Default:
acc_dtype: torch.dtype- Accumulator dtype. Only
torch.float32is supported
- Accumulator dtype. Only
mma_tiler_mn: Tuple[int, int]- Kernel tile size
(TILE_M, TILE_N) TILE_M ∈ {128, 256}TILE_N ∈ {64, 128, 192, 256}
- Kernel tile size
cluster_shape_mn: Tuple[int, int] | None- Thread-block cluster shape
- Default:
(2, 1)whenTILE_M == 256, else(1, 1)
sf_vec_size: int- Scale-factor vector size. Allowed values:
{16, 32}
- Scale-factor vector size. Allowed values:
vector_f32: bool- Enables vectorized f32 operations for supported configurations
- CUDA stream (
current_streamin class API,streamin wrapper)
Wrapper return values
Returns a TupleDict with keys:
c_tensord_tensoramax_tensorsfd_tensor
Tuple unpacking order is: (c_tensor, d_tensor, amax_tensor, sfd_tensor).
Support surface and constraints
Layouts
Amay bem-major ork-majorBmay ben-major ork-majorCandDmust share the same layout- The wrapper exposes this as
c_major ∈ {"m", "n"}
Dtypes
AandBmust have the same dtypeSFA,SFB, andSFDmust have the same dtypesf_vec_size == 32is unsupported withsf_dtype == float8_e4m3fn- FP8 input requires
sf_vec_size == 32 - FP4 input with FP8
Dis unsupported - FP8
Drequires bothSFDandnorm_const_tensor
Environment
- Requires CUDA with SM100+ compute capability
Usage examples
For end-to-end usage and regression coverage, see:
test/python/fe_api/gemm/test_gemm_srelu.pytest/python/fe_api/gemm/test_gemm_srelu_utils.py