Grouped GEMM + GLU + Hadamard (SM100)
Grouped GEMM + GLU + Hadamard (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 fusion: A contiguous grouped block-scaled GEMM fused with a GLU epilogue, a 16-wide Hadamard transform for post-RHT amax computation, and per-expert amax reductions 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.
This kernel performs:
- Block-scaled grouped GEMM over contiguous expert ranges
- GLU epilogue using per-row
probwith SwiGLU, GeGLU, SiTU-GLU, or SReLU - Hadamard transform over 16-token groups of the post-GLU output
- Per-expert amax reductions before and after the Hadamard transform
Shapes
-
Inputs
A: contiguous activation tensor across all groups, shape(valid_m, K, 1)B: weight tensor across all groups, shape(N, K, L)SFA: 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)padded_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)Hadamard: fixed transform matrix, shape(16, 16)
-
Outputs
C: intermediate GEMM result before GLU/Hadamard, shape(valid_m, N, 1)D: activation output before the Hadamard transform, shape(valid_m, N / 2, 1)for GLU activations and(valid_m, N, 1)for SReLUAmax: per-expert amax ofD, shape(L, 1)whenDis fp16/bf16PostRhtAmax: per-expert amax after the normalized Hadamard transform, shape(L, 1)whenDis fp16/bf16
L is the expert count and valid_m = padded_offsets[-1].
Equations
For rows belonging to expert g:
Split the N dimension into consecutive 32-column gate/up blocks:
For SwiGLU (act_func="swiglu"):
For GeGLU (act_func="geglu"):
For SiTU-GLU (act_func="situglu"):
Here beta_1 = situ_beta1 and beta_2 = situ_beta2, with defaults 4.0 and 25.0. situ_beta1 specializes the compiled kernel and is part of its cache key. situ_beta2 is a runtime FP32 scalar, so changing it does not create a new compiled-kernel cache entry.
The returned D is X. For NVFP4 quantization, the kernel also applies the normalized fixed Hadamard matrix H of size 16 x 16 over 16-token groups within each expert and reduces its absolute maximum:
When D is fp16/bf16, the kernel emits both Amax, computed from the untransformed D, and PostRhtAmax. The transformed values are not materialized as another output tensor; the post-RHT amax is intended for the downstream NVFP4 quantization step.
Diagram
API Usage
High-level wrapper
The wrapper constructs the fixed Hadamard matrix internally.
Class API
You may optionally pass a custom sample_hadamard / hadamard_tensor, but the API normalizes it to the fixed 16 x 16 bf16 contiguous layout expected by the kernel. If you do not provide one, the default is the fixed kernel matrix.
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}
- 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 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:
- 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:
- Input tensor Hadamard (optional in class API)
- Shape:
(16, 16) - Dtype:
bfloat16 - Layout: normalized to a contiguous
16 x 16bf16 tensor before compile/execute
- Shape:
- Output tensor C:
result["c_tensor"](wrapper) orsample_c/c_tensor(class)- Shape:
(valid_m, N, 1) - Layout: must be
n-major - Dtype:
{float16, bfloat16}
- Shape:
- Output tensor D:
result["d_tensor"](wrapper) orsample_d/d_tensor(class)- Shape:
(valid_m, N / 2, 1)for GLU activations;(valid_m, N, 1)for SReLU - Layout: must be
n-major - Dtype:
{float16, bfloat16}
- Shape:
- Output tensor Amax:
result["amax_tensor"](wrapper) orsample_amax/amax_tensor(class)- Shape:
(L, 1) - Dtype:
float32
- Shape:
- Output tensor PostRhtAmax:
result["post_rht_amax_tensor"](wrapper) orsample_post_rht_amax/post_rht_amax_tensor(class)- Shape:
(L, 1) - Dtype:
float32 - Semantics: per-expert amax after normalized RHT(16), for downstream NVFP4 quantization
- Shape:
Common parameters
acc_dtype: torch.dtype- Only
torch.float32is supported
- Only
mma_tiler_mn: Tuple[int, int]- Must be
(256, 256)
- Must be
cluster_shape_mn: Tuple[int, int] | None- Default:
(2, 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
act_func: str- Allowed values:
{"swiglu", "geglu", "situglu", "srelu"}
- Allowed values:
situ_beta1: float- Positive finite gate tanh scale for SiTU-GLU; default
4.0 - Compile-time specialized and included in the wrapper cache key
- Positive finite gate tanh scale for SiTU-GLU; default
situ_beta2: float- Positive finite up-branch tanh scale for SiTU-GLU; default
25.0 - Runtime FP32 scalar; changing it reuses the compiled beta1 specialization
- Positive finite up-branch tanh scale for SiTU-GLU; default
- CUDA stream (
current_streamin class API and wrapper)
Wrapper return values
Returns a TupleDict with keys:
c_tensord_tensoramax_tensorpost_rht_amax_tensor
Tuple unpacking order is: (c_tensor, d_tensor, amax_tensor, post_rht_amax_tensor).
Support surface and constraints
- Only dense contiguous grouped weights are exposed in this frontend integration.
- The wrapper constructs the fixed Hadamard matrix internally.
AandBmust be fp4 input tensors.Dis currently supported for{float16, bfloat16}.Nmust be divisible by64.N / 2must be divisible by16.m_alignedmust be256.expert_cntmust be<= 1024.- The kernel requires SM100+.