Grouped GEMM + sReLU (SM100)
Grouped GEMM + sReLU (SM100)
This is an experimental API and subject to change.
JAX support
JAX arrays are not supported: both dense and discrete modes consume the SFA scale-factor tensor as an MMA-permuted strided cute tensor argument, a layout with no row-major (JAX) equivalent. JAX inputs raise a clear ValueError at the entry points. The API is otherwise type-erased and torch-lazy.
Overview
Grouped GEMM + sReLU fusion: A grouped block-scaled GEMM fused with a probability-gated squared-ReLU epilogue on NVIDIA Blackwell GPUs (SM100+), designed for MoE-style workloads. The API supports dense contiguous weights and discrete per-expert weight allocations. Groups are contiguous in the M dimension and described by padded_offsets.
This kernel performs:
- Block-scaled grouped GEMM over contiguous expert ranges
- sReLU epilogue using per-row
prob - Optional output quantization through
SFD_row/SFD_colorAmax
Shapes
-
Inputs
A: contiguous activation tensor across all groups, shape(valid_m, K, 1)B: dense weight tensor across all groups, shape(N, K, L), or discrete per-expert tensors addressed byb_ptrsSFA: shape(32, 4, ceil_div(valid_m, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), 1)SFB: dense scale-factor tensor, shape(32, 4, ceil_div(N, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), L), or discrete per-expert tensors addressed bysfb_ptrspadded_offsets: cumulative padded group ends, shape(L,)alpha: per-group scaling factors, shape(L,)prob: per-row gating probabilities, shape(valid_m, 1, 1)
-
Outputs
C: intermediate GEMM result, shape(valid_m, N, 1)D: row output after sReLU, shape(valid_m, N, 1)D_col: column output after sReLU, shape(valid_m, N, 1)SFD_row: shape(32, 4, ceil_div(valid_m, 128), 4, ceil_div(ceil_div(N, sf_vec_size), 4), 1)whenDis FP8SFD_col: shape(32, 4, ceil_div(N, 128), 4, ceil_div(ceil_div(valid_m, sf_vec_size), 4), 1)whenDis FP8Amax: shape(L, 1)whenDis fp16/bf16
L is the expert count and valid_m = padded_offsets[-1].
Equations
For rows belonging to expert g:
Tanh soft clamp (tanh_clamp_scale)
Passing tanh_clamp_scale=s (a positive float; None, the default, keeps the equation above)
soft-clamps the ReLU output before squaring, bounding D by :
s is baked into the compiled kernel, so distinct values compile distinct kernels — it is meant
to be constant for a whole training job, not a per-call argument.
The epilogue is evaluated in FP32 with the approximate tanh. Measured on sm_103 against a
correctly-rounded reference, over a sweep of plus saturated inputs, its maximum
absolute error is (the exact variant measures on the same
sweep). Since with , that propagates to at most
absolutely — about at — and to roughly
relative in the saturated tail. is capped at 1 before use, so and the output bound
hold structurally rather than by relying on the range of the
approximate instruction.
When D is FP8, even that small error is enough to carry a value across an output-format
rounding boundary, so individual elements can land one ULP away from a reference computed with
an exact tanh (12.5% relative for e4m3). Measured incidence on this kernel’s test matrix is
a few elements in . This does not affect the unclamped path, where the kernel and
a reference both evaluate and therefore quantize identically.
The matching backward kernel takes the same tanh_clamp_scale — see
grouped_gemm_dsrelu. Both must be built with the same s, or the saved
pre-activation is differentiated against the wrong nonlinearity.
D_col stores the same logical output in the column-quantized path used by the grouped kernel family. When D is FP8, the kernel also emits SFD_row and SFD_col. When D is fp16/bf16, the kernel can emit per-expert Amax.
Diagram
API Usage
High-level wrapper
Discrete-weight wrapper
Class API
Parameters
Input/Output tensors
- Input tensor A:
a_tensor(wrapper) orsample_a/a_tensor(class)- Shape:
(valid_m, K, 1) - Layout: must be
k-major - 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) - Layout: must be
k-major - Dtype: must match
A
- Shape:
- Discrete input B pointers:
b_ptrs(wrapper) ornum_experts/b_shape/b_dtype(class)b_ptrs: 1-Dint64CUDA tensor containing one data pointer per expertnandb_dtypeare required in wrapper discrete modeb_majormay be"k"or"n"for supported FP8 cases; FP4 uses"k"
- Input tensor SFA:
sfa_tensor(wrapper) orsample_sfa/sfa_tensor(class)- Shape:
(32, 4, ceil_div(valid_m, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), 1) - 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:
- Discrete input SFB pointers:
sfb_ptrs- 1-D
int64CUDA tensor containing one scale-factor pointer per expert
- 1-D
- Input tensor padded_offsets
- Shape:
(L,) - Dtype:
int32
- Shape:
- Input tensor alpha
- Shape:
(L,) - Dtype:
float32
- Shape:
- Input tensor prob
- Shape:
(valid_m, 1, 1) - Dtype:
float32 - Required
- Shape:
- Output tensor C:
result["c_tensor"](wrapper) orsample_c/c_tensor(class)- Shape:
(valid_m, N, 1) - Layout: must be
n-major - Dtype:
{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}
- Shape:
- Output tensor D:
result["d_tensor"](wrapper) orsample_d/d_tensor(class)- Shape:
(valid_m, N, 1) - Layout: must be
n-major - Dtype:
- FP4 input:
{float16, bfloat16, float32} - FP8 input:
{float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}
- FP4 input:
- Shape:
- Output tensor D_col:
result["d_col_tensor"](wrapper) orsample_d_col/d_col_tensor(class)- Shape:
(valid_m, N, 1) - Layout: must match
D - Dtype: must match
D
- Shape:
- Output tensors SFD_row / SFD_col
- Dtypes: must match
SFA - Required when FP8 output scale factors are generated
- Dtypes: must match
- Output tensor Amax
- Shape:
(L, 1) - Dtype:
float32
- Shape:
- Input tensor Norm Const
- Shape:
(1,) - Dtype:
float32 - Required when FP8 output scale factors are generated
- Shape:
Common parameters
acc_dtype: torch.dtype- Only
torch.float32is supported
- Only
mma_tiler_mn: Tuple[int, int]TILE_Mdepends on the 1-CTA / 2-CTA modeTILE_N ∈ {128, 256}
cluster_shape_mn: Tuple[int, int] | None- Default:
(2, 1)whenTILE_M == 256, else(1, 1)
- Default:
sf_vec_size: int- Allowed values:
{16, 32}
- Allowed values:
vector_f32: bool- Enables vectorized f32 operations for supported configurations
m_aligned: int- Must equal the kernel fixed pad size
256
- Must equal the kernel fixed pad size
discrete_col_sfd: bool- Enables the discrete column-scale-factor path used by grouped FP8
- CUDA stream (
current_streamin class API,current_streamin wrapper)
Wrapper return values
Returns a TupleDict with keys:
c_tensord_tensord_col_tensoramax_tensorsfd_row_tensorsfd_col_tensor
Tuple unpacking order is: (c_tensor, d_tensor, d_col_tensor, amax_tensor, sfd_row_tensor, sfd_col_tensor).
Support surface and constraints
Layouts
Amust bek-majorBmust bek-major- Discrete
Bsupportsb_major="k"and supported FP8b_major="n"configurations C,D, andD_colmust ben-major- The wrapper only supports
cd_major="n"
Dtypes
AandBmust have the same dtypeSFA,SFB,SFD_row, andSFD_colmust have the same dtype- Scale-factor dtype constraint:
sf_vec_size == 32is unsupported whensf_dtype == float8_e4m3fn - Input dtype constraint: FP8
A/Binputs requiresf_vec_size == 32 - Grouped FP8 currently requires
discrete_col_sfd=True
Shapes and environment
prob_tensoris requiredm_alignedmust be256- Requires CUDA with SM100+ compute capability
Usage examples
For end-to-end usage and regression coverage, see:
test/python/fe_api/grouped_gemm/test_grouped_gemm_srelu.pytest/python/fe_api/grouped_gemm/test_grouped_gemm_srelu_utils.py