Grouped GEMM + GLU + Hadamard + Quant (SM100)
Grouped GEMM + GLU + Hadamard + Quant (SM100)
This is an experimental API and subject to change.
JAX support
JAX arrays are not supported: this fusion is block-scaled-only and its mandatory scale-factor inputs use an MMA-interleaved 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 + GLU + Hadamard + Quant fusion: A contiguous grouped block-scaled GEMM fused with a GLU/SReLU epilogue, optional RHT (Hadamard transform) output, and optional NVFP4 output quantization on NVIDIA Blackwell GPUs (SM100+), designed for MoE-style workloads. Groups are contiguous in the M dimension and described by padded_offsets.
This frontend integration is currently wired for the FP4 input path and exposes the quantized Hadamard fusion under the operation name:
GroupedGemmGluHadamardQuantSm100grouped_gemm_glu_hadamard_quant_wrapper_sm100
This kernel performs:
- Block-scaled grouped GEMM over contiguous expert ranges
- GLU or SReLU epilogue using per-row
prob - Optional NVFP4 quantization of the post-activation
Doutput - Optional RHT output in bf16 or NVFP4 form
Shapes
Let N_out = N / 2 for act_func="swiglu" or "geglu" and N_out = N for act_func="srelu". Let SF(rows, cols) denote the swizzled scale-factor layout (32, 4, ceil_div(rows, 128), 4, ceil_div(ceil_div(cols, sf_vec_size), 4), 1).
-
Inputs
A: contiguous activation tensor across all groups, shape(valid_m, K, 1)B: weight tensor across all groups, shape(N, K, L)in dense modeb_ptrs/sfb_ptrs: int64 CUDA pointer arrays, shape(L,), in discrete modeSFA: shape(32, 4, ceil_div(valid_m, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), 1)SFB: shape(32, 4, ceil_div(N, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), L)in dense modepadded_offsets: cumulative padded group ends, shape(L,)alpha: per-group scaling factors, shape(L,)prob: per-row gating probabilities, shape(valid_m, 1, 1)bias(optional): per-expert bias tensor, shape(N, L)with stride(1, N)
-
Outputs
C: intermediate GEMM result before activation/clamping, shape(valid_m, N, 1)D: post-activation output, logical shape(valid_m, N_out, 1)SFD: swizzled e4m3 scale factors for NVFP4D, shapeSF(valid_m, N_out), present only whenDis NVFP4RHT rowwise: optional Hadamard-transform output across feature blocks, logical shape(valid_m, N_out, 1)SFRHT rowwise: scale factors for NVFP4 rowwiseRHT, shapeSF(valid_m, N_out), present only when rowwiseRHTis NVFP4RHT colwise: optional NVFP4 Hadamard-transform output across token blocks, flattened in per-expert transposed orderSFRHT colwise: scale factors for NVFP4 colwiseRHT, flattened in per-expert ragged order
For packed NVFP4 output tensors (torch.float4_e2m1fn_x2), the physical tensor stores two logical values per byte along the innermost dimension. The wrapper therefore allocates packed D and rowwise RHT tensors with physical second dimension N_out / 2. Raw torch.uint8 tensors are not accepted as a packed FP4 container by this fusion.
Colwise RHT is NVFP4-only. Its data tensor is a flat physical tensor of length valid_m * N_out / 2; after FP4 unpacking, each expert segment is logical (N_out, expert_m) with adjacent token values packed together. Expert segments are concatenated in expert order. SFRHT colwise is flat length valid_m * N_out / sf_vec_size; each expert segment is SF(N_out, expert_m) flattened in expert order, so the scale segment for an expert starts at N_out * expert_m_prefix / sf_vec_size.
L is the expert count. valid_m is the M extent of a_tensor; by contract it must match the final cumulative padded offset in padded_offsets.
Equations
For rows belonging to expert g:
The bias term is omitted when bias_tensor=None. C stores this unclamped GEMM-plus-bias value.
Split the N dimension into consecutive 32-column gate/up blocks:
When glu_limit is set for swiglu or geglu, both G_b and U_b are clamped to [-glu_limit, glu_limit] before activation. Let gamma be glu_alpha when it is set and not 1.0; otherwise gamma = 1.
For SwiGLU (act_func="swiglu"):
For GeGLU (act_func="geglu"):
For SReLU (act_func="srelu"), the kernel does not split N into gate/up halves:
When requested, the RHT output applies a fixed 16-wide orthonormal Hadamard transform to bf16-rounded D. Rowwise RHT transforms feature blocks. Colwise RHT transforms token blocks and writes NVFP4 data in per-expert transposed order.
When D or RHT is NVFP4, the kernel emits packed e2m1 data plus e4m3 scale factors. norm_const and rht_norm_const are the corresponding global encode scales.
Diagram
API Usage
High-level wrapper
Leave both rht_rowwise_dtype and rht_colwise_dtype unset to skip the Hadamard/RHT output. Set rht_rowwise_dtype=torch.bfloat16 for unquantized rowwise RHT, rht_rowwise_dtype=torch.float4_e2m1fn_x2 for quantized rowwise RHT, or rht_colwise_dtype=torch.float4_e2m1fn_x2 for wgrad-compatible colwise RHT.
Class API
Discrete weight mode
The wrapper also accepts per-expert discrete weight allocations:
b_ptrs and sfb_ptrs must be contiguous int64 CUDA tensors containing device pointers for each expert. Each b_ptrs entry points to a logical (N, K) FP4 expert weight allocation; each sfb_ptrs entry points to that expert’s scale-factor allocation with logical shape (32, 4, ceil_div(N, 128), 4, ceil_div(ceil_div(K, sf_vec_size), 4), 1).
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
- 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:
- Input tensor B pointers:
b_ptrs(wrapper execute) ornum_experts/b_shape/b_dtype(class construction)- Shape:
(L,) - Dtype:
int64 - Device: CUDA
b_dtypemust matchA; FP4 discrete mode requiresb_major="k"
- Shape:
- 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} - Set
sf_fp8_dtype_override="e5m3"to reinterpretfloat8_e4m3fnstorage as UE5M3 on Rubin.
- 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 SFB pointers:
sfb_ptrsin discrete mode- Shape:
(L,) - Dtype:
int64 - Device: CUDA
- Shape:
- 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
- Shape:
- Input tensor bias (optional)
- Shape:
(N, L) - Stride:
(1, N) - Dtype:
{float16, bfloat16, float32}
- Shape:
- Output tensor C
- Shape:
(valid_m, N, 1) - Layout: must be
n-major - Dtype:
{float16, bfloat16}
- Shape:
- Output tensor D
- Logical shape:
(valid_m, N_out, 1) - Layout: must be
n-major - Dtype:
{bfloat16, float4_e2m1fn_x2} - NVFP4
DrequiresSFD
- Logical shape:
- Output tensor SFD (present only with NVFP4
D)- Shape:
SF(valid_m, N_out)=(32, 4, ceil_div(valid_m, 128), 4, ceil_div(ceil_div(N_out, sf_vec_size), 4), 1) - Layout: swizzled scale-factor layout matching
SFA - Dtype:
float8_e4m3fn
- Shape:
- Output tensor RHT rowwise (optional)
- Logical shape:
(valid_m, N_out, 1) - Layout: must be
n-major - Dtype:
{bfloat16, float4_e2m1fn_x2} - NVFP4 rowwise
RHTrequiresSFRHT rowwise
- Logical shape:
- Output tensor SFRHT rowwise (present only with NVFP4 rowwise
RHT)- Shape:
SF(valid_m, N_out) - Layout: swizzled scale-factor layout
- Dtype:
float8_e4m3fn, independent ofSFA/SFBdtype
- Shape:
- Output tensor RHT colwise (optional)
- Logical unpacked shape: per-expert
(N_out, expert_m)segments concatenated in expert order - Physical shape: 1-D packed FP4 tensor with
valid_m * N_out / 2elements - Dtype:
float4_e2m1fn_x2 - NVFP4 colwise
RHTalways requiresSFRHT colwise
- Logical unpacked shape: per-expert
- Output tensor SFRHT colwise (present with colwise
RHT)- Shape:
(valid_m * N_out / sf_vec_size,) - Layout: flattened per-expert
SF(N_out, expert_m)segments in expert order - Dtype:
float8_e4m3fn, independent ofSFA/SFBdtype
- Shape:
Configuration
act_func:"swiglu","geglu", or"srelu"cd_major: must be"n"mma_tiler_mn: must be(256, 256)cluster_shape_mn: cluster dimensions; for this fixed 2-CTA tiler,cluster_shape_mn[0]must be2andcluster_shape_mn[1]must be a positive power of two no larger than4sf_vec_size: must be16sf_fp8_dtype_override:Noneuses the scale format implied bySFA/SFBdtype."e5m3"reinterpretstorch.float8_e4m3fnSFA/SFB storage as UE5M3 input scale factors; this is Rubin-only and does not convert tensor contents.m_aligned: must be256rht_rowwise_dtype: unset to skip rowwise RHT,torch.bfloat16for unquantized rowwise RHT, ortorch.float4_e2m1fn_x2for NVFP4 rowwise RHTrht_colwise_dtype: unset to skip colwise RHT, ortorch.float4_e2m1fn_x2for NVFP4 colwise RHT in per-expert transposed orderglu_alpha: optional final output scale forswiglu/gegluglu_limit: optional clamp limit applied to both gate and up blocks forswiglu/geglunorm_const: global encode scale for NVFP4Drht_norm_const: global encode scale for NVFP4RHT
Constraints
- Requires SM100 or newer.
Nmust be divisible by64.Kmust satisfy the FP4 K-major 16-byte alignment requirement, soKmust be divisible by32.valid_m/a_tensor.shape[0]must be divisible by256;padded_offsetsmust contain cumulative 256-aligned expert ends and end atvalid_m.N_outmust be divisible by32.- NVFP4 quantization requires
N_outdivisible by128. - NVFP4 quantization is not supported with
act_func="srelu". sf_fp8_dtype_override="e5m3"requires Rubin (SM107) andSFA/SFBtensors stored astorch.float8_e4m3fn.- At most one of rowwise and colwise RHT may be requested.
- Colwise RHT is NVFP4-only and uses flat per-expert ragged scale storage.
expert_cntmust be<= 1024.- Dense and discrete weight modes are mutually exclusive.