Grouped GEMM + GLU (SM100)
Grouped GEMM + GLU (SM100)
This is an experimental API and subject to change.
JAX support
Supports JAX arrays on the BF16 backend in discrete weight mode (swiglu and geglu): b_ptrs as a packed little-endian uint8 pointer array (8 bytes per pointer; int64 accepted with jax x64 mode), outputs allocated as n-major C-contiguous jnp arrays. Dense b_tensor (expert-outermost strides), column-major bias_tensor, and the block-scaled backend (MMA-interleaved scale-factor layouts) are not expressible as JAX arrays and raise clear errors. The wrapper is eager, on the CUDA legacy default stream: block_until_ready inputs, synchronize before reading outputs; keep weight arrays alive until the kernel completes.
For jitted JAX programs use the jax.jit-compatible XLA custom-call entry point grouped_gemm_glu_jax_sm100 (built on cudnn.jax.call; discrete mode, no bias, b_major="k"): outputs are fresh XLA-managed arrays with rows at/past padded_offsets[-1] zero-filled, no manual synchronization needed. linear_offset is a compile-time constant (each distinct value compiles a new specialization). Under tracing the padded_offsets values cannot be host-validated, and the per-expert weight buffers behind b_ptrs must stay alive and unmoved across every execution of the traced computation.
Overview
Unified Grouped GEMM + GLU fusion: one public class and wrapper select a plain BF16 or legacy block-scaled grouped GEMM fused with a GLU epilogue (SwiGLU, GeGLU, or block-scaled SiTU-GLU) on NVIDIA Blackwell GPUs (SM100/SM103), with a block-scaled forward path on Rubin (SM107). The operation is implemented with CUTLASS/CuTe DSL.
This is a unified API that supports both weight layout modes:
- Dense mode: All expert weights packed into a single contiguous
(N, K, L)tensor - Discrete mode: Per-expert weight pointers (no weight stacking required)
Supported activation functions:
- SwiGLU:
act_func="swiglu"(default) - GeGLU:
act_func="geglu" - SiTU-GLU:
act_func="situglu"(block-scaled SM100/SM103 only)
Groups are contiguous in the M dimension and described by padded_offsets (cumulative aligned end offsets).
Backend dispatch
Mixed families and unsupported pairs are rejected before allocation or compilation. Each backend’s argument contract is described below. SiTU-GLU is not available on the BF16 or Rubin backends. The block-scaled GeGLU forward path supports runtime activation alpha, clamp limits, and linear offset on both Blackwell and Rubin. The corresponding dGeGLU backward path supports the same activation configuration, with the parameter values included in the backward compiled-kernel cache key.
BF16 contract
Pass sfa_tensor=None, sfb_tensor=None (or sfb_ptrs=None), and
norm_const_tensor=None; keep sf_vec_size=16 and discrete_col_sfd=False.
Non-None scale controls are an error.
Tensors, layouts, and equation
For padded rows M, reduction dimension K, pre-GLU width N, and L
experts, BF16 uses:
A:(M, K, 1), stride(K, 1, M*K), BF16;- dense
B:(N, K, L), K-major stride(K, 1, N*K), BF16; - discrete
b_ptrs: contiguous CUDA int64 pointers to expert(N, K)BF16 matrices, withn=N,b_dtype=torch.bfloat16, andb_major="k"or"n"; padded_offsets:(L,), stride(1,), int32 cumulative 256-aligned ends;alpha:(L,), FP32;prob:(M, 1, 1), stride(1, 1, 1), FP32;- optional bias:
(N, L), stride(1, N), BF16/FP16/FP32; C:(M, N, 1), stride(N, 1, M*N);D:(M, N/2, 1), stride(N/2, 1, M*N/2).
For expert g, first compute
Columns are paired as alternating 32-wide gate/up blocks. For SwiGLU,
For GeGLU, let gate = min(gate(C), 7),
up = clamp(up(C), -7, 7), geglu_alpha=1.702, and the default
linear_offset=1:
C/D may be BF16, FP16, or FP32. N is divisible by 64. The pointer-array
tensor is stream-recorded; every pointed allocation must remain alive and
unchanged until the launch stream completes.
The wrapper return order is exactly c_tensor, d_tensor, d_col_tensor,
amax_tensor, sfd_row_tensor, sfd_col_tensor. On BF16,
d_col_tensor, amax_tensor, sfd_row_tensor, and sfd_col_tensor are
always None; c_tensor is None unless generate_c=True.
Block-scaled contract
The block-scaled backend performs:
- Block-scaled grouped GEMM: Low-precision GEMM (FP4, FP8) with per-block scale factors across multiple expert groups
- GLU activation: Fused SwiGLU, GeGLU, or SiTU-GLU activation applied to the GEMM output
- Optional quantized output: Produces row and column scale factors for downstream quantization
Shapes
Equations
For SiTU-GLU, with gate branch G and up branch U, the fused epilogue computes
where beta_1 = situ_beta1 and beta_2 = situ_beta2, with defaults
beta_1 = 4.0 and beta_2 = 25.0. situ_beta1 specializes the compiled kernel
and is part of its cache key; situ_beta2 is a runtime FP32 scalar and does not
create a new compiled-kernel cache entry.
- Inputs
A: contiguous activation tensor across all groups, shape(valid_m, K, 1)B(dense): weight tensor across all groups, shape(N, K, L)B(discrete): per-expert weight pointers,b_ptrsshape(num_experts,)of int64SFA: scale factor tensor for A, shape(32, 4, ceil(valid_m/128), 4, ceil(ceil(K/sf_vec_size)/4), 1)SFB(dense): scale factor tensor for B, shape(32, 4, ceil(N/128), 4, ceil(ceil(K/sf_vec_size)/4), L)SFB(discrete): per-expert SFB pointers,sfb_ptrsshape(num_experts,)of int64padded_offsets: cumulative sum of aligned group M sizes, shape(L,).valid_m = padded_offsets[-1]alpha: per-group scaling factors, shape(L,)bias(optional): per-expert bias tensor, shape(N, L)with stride(1, N)prob: per-row gating probabilities, shape(valid_m, 1, 1)norm_const: normalization constant for FP8 quantization, shape(1,)
- Outputs
C: intermediate GEMM result, shape(valid_m, N, 1)D: row-quantized GLU output, shape(valid_m, N/2, 1)D_col: column-quantized GLU output, shape(valid_m, N/2, 1)SFD_row: row scale factors (whend_dtypeis FP8), shape(32, 4, ceil(valid_m/128), 4, ceil(ceil((N/2)/sf_vec_size)/4), 1)SFD_col: column scale factors (whend_dtypeis FP8), shape(32, 4, ceil((N/2)/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1)amax: per-group amax (whend_dtypeis bf16/fp16), shape(L, 1)
Step 1: Block-scaled grouped GEMM (per group g with rows m in [padded_offsets[g-1], padded_offsets[g])):
Step 2: GLU epilogue (performed by pairing 32-column blocks along N):
Let block size G = 32. For each pair of consecutive 32-wide column blocks:
- Gate block:
G_b = C[:, 2·b·G : 2·b·G + G] - Up block:
U_b = C[:, 2·b·G + G : 2·b·G + 2·G]
For SwiGLU (act_func="swiglu"):
For GeGLU (act_func="geglu"):
Only the gate’s upper bound is clamped. The offset is added after clamping
the up branch. The gate nonlinearity is g * sigmoid(geglu_alpha * g);
silu(geglu_alpha * g) would introduce an extra factor of geglu_alpha.
The optional C output stores the GEMM result before clamping. The epilogue
uses the FP32 accumulator, without rounding it to the C output dtype first.
The defaults are geglu_alpha=1.702, glu_clamp_max=7.0,
glu_clamp_min=-7.0, and linear_offset=1.0. For DeepSeek V4 clamped SwiGLU,
select act_func="geglu" with geglu_alpha=1.0, linear_offset=0.0,
glu_clamp_max=L, and glu_clamp_min=-L, where L is the model’s clamp
limit. act_func="swiglu" does not apply these clamp parameters.
For block-scaled dense and discrete calls, these four activation parameters
are runtime FP32 scalars: changing their values reuses the same compiled
kernel. geglu_alpha scales the sigmoid input and is independent of the
per-expert GEMM scaling tensor alpha_tensor.
Step 3: Optional output quantization (when SFD outputs are generated):
Diagram
API usage
BF16
High-level wrapper
Class API
use_dynamic_sched=False uses static scheduling; True caches a dynamic-M
callable for compatible shapes. Cache keys include compile-sensitive layouts,
dtypes, features, activation, scheduler, tile/cluster, output policy, and
overlap margin, but not the runtime GeGLU linear_offset.
Block-scaled
High-level wrapper
Dense mode:
Discrete mode:
bias_tensor must use the kernel layout expected by the fused bias path: shape (N, L) and stride (1, N).
Class API
Dense mode:
sample_bias and runtime bias_tensor must both use shape (N, L) and stride (1, N).
Discrete mode:
Parameters
Weight Mode
The weight mode is auto-detected from constructor arguments:
- Dense: Provide
sample_bandsample_sfb(contiguous weight tensors) - Discrete: Provide
num_experts,b_shape, andb_dtype(per-expert pointer mode)
Providing both or neither raises ValueError.
Input/Output Tensors
-
Input tensor A:
a_tensor/sample_a- Shape:
(valid_m, K, 1) - Stride:
(K, 1, valid_m·K)— must be K-major - Dtype (
ab_dtype):{float4_e2m1fn_x2, uint8, float8_e4m3fn, float8_e5m2}
- Shape:
-
Input tensor B (dense mode):
b_tensor/sample_b- Shape:
(N, K, L)whereL = num_groups - Stride:
(K, 1, N·K)— must be K-major - Dtype: Must match A
- Shape:
-
Input B pointers (discrete mode):
b_ptrs- Shape:
(num_experts,)— 1-D int64 device tensor of per-expert B data pointers - Build via:
torch.tensor([b.data_ptr() for b in experts], dtype=torch.int64, device="cuda")
- Shape:
-
Output tensor C: returned in wrapper dict or
c_tensorin class- Shape:
(valid_m, N, 1) - Stride:
(N, 1, valid_m·N)— must be N-major - Dtype (
c_dtype):{float16, bfloat16}for FP4 inputs;{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}otherwise
- Shape:
-
Output tensor D:
d_tensor/sample_d- Shape:
(valid_m, N/2, 1) - Stride:
(N/2, 1, valid_m·N/2)— must be N-major - Dtype (
d_dtype):{bfloat16, float32}for FP4 inputs;{float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}otherwise
- Shape:
-
Output tensor D_col:
d_col_tensor/sample_d_col- Shape:
(valid_m, N/2, 1)— must match D dtype and stride
- Shape:
-
Scale factor tensors: Same as contiguous swiglu (SFA, SFB, SFD_row, SFD_col)
- SFB (discrete mode): use
sfb_ptrs(1-D int64 device tensor of per-expert SFB pointers)
- SFB (discrete mode): use
-
Group offsets:
padded_offsets— shape(L,), dtypeint32 -
Scaling tensors:
alphashape(L,),probshape(valid_m, 1, 1),amaxshape(L, 1),norm_constshape(1,)
Common Parameters
acc_dtype: Must betorch.float32mma_tiler_mn: Kernel tile size(TILE_M, TILE_N). Default:(256, 256)TILE_M ∈ {128, 256}TILE_N = 256
cluster_shape_mn: Thread Block cluster shape. Default:(2, 1)whenTILE_M=256,(1, 1)otherwisesf_vec_size: Scale factor vector size.{16, 32}. Default:16vector_f32: Enable packed f32 operations. Default:Falsem_aligned: Must be256(FIX_PAD_SIZE). Default:256discrete_col_sfd: Generate discrete col-major scale factors. Default:Falseact_func: Activation function."swiglu"(default),"geglu", or block-scaled"situglu"linear_offset: Offset added to the clamped GeGLU up branch. Default:1.0for GeGLUgeglu_alpha: GeGLU sigmoid input scale. Default:1.702glu_clamp_max: GeGLU upper bound for gate and up. Default:7.0glu_clamp_min: GeGLU lower bound for up only. Default:-7.0- The activation alpha and clamp controls above are supported by the block-scaled forward backend on Blackwell and Rubin; the BF16 contract retains its fixed values
situ_beta1: Positive finite gate tanh scale for SiTU-GLU. Default:4.0situ_beta2: Positive finite up-branch tanh scale for SiTU-GLU. Default:25.0b_major(discrete only): B tensor major dimension."k"(default) or"n". Must be"k"for FP4.
Wrapper-specific Parameters
c_dtype: Intermediate C tensor data type. Default:torch.bfloat16d_dtype: Output D tensor data type. Default:torch.bfloat16cd_major: Must be"n". Default:"n"n(discrete only): B weight N dimension (full N before GLU split)b_dtype(discrete only): B weight data type
Wrapper Return Values
Returns a TupleDict (dictionary + tuple unpacking):
c_tensor: Intermediate GEMM resultd_tensor: Row-quantized GLU outputd_col_tensor: Column-quantized GLU outputamax_tensor: Per-group amax (whend_dtypeis bf16/fp16)sfd_row_tensor: Row scale factors (when SFD enabled)sfd_col_tensor: Column scale factors (when SFD enabled)
Support Surface and Constraints
Layouts and Strides
Amust be K-majorBmust be K-major (dense mode). For discrete mode: K-major or N-major (K-major required for FP4)C,D,D_colmust be N-major- All tensors must be 16-byte aligned along the contiguous dimension
Data Types
Additional Type Constraints
AandBmust have the same dtype- Scale factor tensors (SFA, SFB, SFD_row, SFD_col) must have the same dtype
DandD_colmust have the same dtypebiasmust be one of{float16, bfloat16, float32}biasmust have shape(N, L)and stride(1, N)- For non-bias paths, FP4
ab_dtypewithsf_vec_size=16andd_dtype=float32is not supported - FP4
ab_dtyperequiresc_dtypein{float16, bfloat16}
Shapes and Divisibility
Nmust be divisible by 64 (two consecutive 32-column blocks for GLU pairing)- Expert count must be
<= 1024 - Each group’s M dimension is aligned to
m_aligned(256) - All supported kernel configurations require
mma_tiler_mn[1] == 256 use_single_group_runtime_offsets=Trueis supported only by the block-scaled kernel with exactly one expert. In this mode the kernel derivespadded_offsets[0]from runtimeA.shape[0]and does not load its value from device memory; the argument must still be an int32 tensor with shape(1,).
Environment
- Requires CUDA with SM100/SM103 (Blackwell), or SM107 (Rubin) for the block-scaled forward backend
- Rubin MXFP8 uses matching FP8 A/B operands, E8M0 scale factors, and
sf_vec_size=32
Usage Examples
For usage examples, see test cases in test/python/fe_api/grouped_gemm/test_grouped_gemm_glu.py (dense mode, unified API) and test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_swiglu.py (discrete mode).