Grouped GEMM + dGLU (SM100)
Grouped GEMM + dGLU (SM100)
This is an experimental API and subject to change.
JAX support
Supports JAX arrays on the BF16 backend in discrete weight mode (dswiglu and dgeglu), including generate_dbias=True and caller-provided zero-initialized dprob. Dense b_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_dglu_jax_sm100 (built on cudnn.jax.call; discrete mode): dprob and (with generate_dbias=True) dbias come back as bridge-managed zero-initialized accumulator outputs — no caller-zeroed buffers, no manual synchronization. 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 + dGLU fusion: one public class and wrapper select a plain BF16 or legacy block-scaled grouped GEMM fused with a dGLU backward epilogue (dSwiGLU, dGeGLU, or block-scaled dSiTU-GLU) on NVIDIA Blackwell GPUs (SM100/SM103), with a block-scaled dSwiGLU/dGeGLU 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)
And both backward activation functions:
- dSwiGLU:
act_func="dswiglu"(default) - dGeGLU:
act_func="dgeglu" - dSiTU-GLU:
act_func="dsituglu"(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.
dsituglu is not available on the BF16 or Rubin backends.
The block-scaled dGeGLU path supports configurable activation alpha, clamp
bounds, and linear offset on both Blackwell and Rubin in dense and discrete
weight modes. These values are constructor configuration and remain part of
the backward wrapper’s 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, discrete_col_sfd=False, and
epilogue_op=None. BF16 also uses the source GeGLU constants
geglu_alpha=1.702, glu_clamp_max=7, and glu_clamp_min=-7.
Tensors, layouts, and equation
For padded rows M, reduction dimension K, compact gradient 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"; C:(M, 2N, 1), stride(2N, 1, 2M*N), BF16/FP16/FP32;padded_offsets:(L,)int32 cumulative 256-aligned ends;alpha,beta:(L,)FP32;prob:(M, 1, 1)FP32;- caller-zeroed
dprob:(M, 1, 1), stride(1, 1, 1), FP32; D_row:(M, 2N, 1), stride(2N, 1, 2M*N), BF16/FP16/FP32;- caller-zeroed optional
dbias:(L, 2N, 1), stride(2N, 1, 1), BF16.
For expert g, the compact GEMM gradient and scaled forward activation are
Split X into alternating 32-wide gate/input blocks. For dSwiGLU, with
s = sigmoid(gate):
For dGeGLU, distinguish raw values, clamped activation values, and the source’s value-bearing filters:
With the default linear_offset=1, the kernel computes
dprob accumulates the row sum of the matching unscaled activation times R;
dbias is the per-expert row reduction of interleaved D_row. The
pointer-array tensor is stream-recorded, while every pointed allocation must
remain alive and unchanged until that stream completes.
The wrapper return order is exactly d_row_tensor, d_col_tensor,
dprob_tensor, dbias_tensor, amax_tensor, sfd_row_tensor,
sfd_col_tensor. On BF16, d_col_tensor, amax_tensor, sfd_row_tensor, and
sfd_col_tensor are None; dbias_tensor is None unless
generate_dbias=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
- dGLU backward epilogue: Fused backward computation using the forward
Ctensor (input/gate interleaved) - Optional quantized output: Produces row and column scale factors for downstream quantization
Shapes
Equations
For dSiTU-GLU, define
The fused backward computes
and returns ref * prob * T_u * dT_g/dG and
ref * prob * T_g * dT_u/dU. dprob accumulates the reduction of
ref * T_g * T_u across the output columns in 32-column chunks, producing
shape (valid_m, 1, 1).
The beta values are compile-time specialization values and therefore belong to
the dGLU compiled-kernel cache key.
- 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 int64C: forward intermediate tensor with interleaved input/gate blocks, shape(valid_m, 2N, 1)SFA: 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 for GEMM, shape(L,)beta: per-group scaling factors forCin dSwiGLU, shape(L,); dGeGLU consumes the savedCdirectly and ignoresbetaprob: per-row gating probabilities, shape(valid_m, 1, 1)norm_const: normalization constant for FP8 quantization, shape(1,)
- Outputs
D_row: row-quantized dGLU output, shape(valid_m, 2N, 1)D_col: column-quantized dGLU output, shape(valid_m, 2N, 1)dprob: gradient ofprob, shape(valid_m, 1, 1). Must be zero-initialized.dbias(optional): per-expert bias gradient tensor, shape(L, 2N, 1)SFD_row: row scale factors (whend_dtypeis FP8), shape(32, 4, ceil(valid_m/128), 4, ceil(ceil((2N)/sf_vec_size)/4), 1)SFD_col: column scale factors (whend_dtypeis FP8), shape(32, 4, ceil((2N)/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1)amax: per-group amax (whend_dtypeis bf16/float16), shape(L, 2, 1)
Step 1: Block-scaled grouped GEMM (per group g with rows m in [padded_offsets[g-1], padded_offsets[g])):
Step 2: dGLU backward epilogue (performed with 32-column interleaving along 2N):
- Deinterleave
Cinto alternating 32-column gate and up blocks. dSwiGLU appliesbeta_gtoC; dGeGLU uses the saved forward values directly.
For dSwiGLU (act_func="dswiglu"):
swish = gate * sigmoid(gate)dprob += sum(swish * input * ref)over 32-column chunksab = ref * prob * swishdswiglu = ref * prob * input * sigmoid(gate) * (1 + gate * (1 - sigmoid(gate)))- Interleave
[ab, dswiglu]back intoD_row/D_colwith 32-column blocks.
For dGeGLU (act_func="dgeglu"), let gate and up be the original
saved C values and ref the GEMM result from Step 1:
The gate has only an upper clamp; the up branch has both bounds. Clamp masks
use the original values before clamping, and gradients are retained at equality
with either bound. The offset is added after clamping the up branch. The
probability gradient does not include a factor of prob, so it can be nonzero
when the routing probability is zero. [dgate, dup] is stored in alternating
32-column blocks.
The upstream ref includes alpha_tensor[g] ** 2 on both architectures; this
per-group GEMM scale is separate from geglu_alpha. dGeGLU reads C in its
stored precision, so a BF16 saved intermediate is the input to this derivative,
not the original FP32 forward accumulator. C is not modified.
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,
use act_func="dgeglu", geglu_alpha=1.0, linear_offset=0.0,
glu_clamp_max=L, and glu_clamp_min=-L, matching the forward configuration
(L=10 for the Flash recipe).
Set these parameters when constructing GroupedGemmDgluSm100 or calling
grouped_gemm_dglu_wrapper_sm100. The backward API specializes the activation
configuration; changing it selects another compiled-kernel cache entry, while
repeating a configuration reuses that entry. It does not expose the forward
API’s per-execution activation controls. With BF16 C, the packed derivative
represents clamp bounds in BF16; limits such as 7 and 10 are exact, while
nonrepresentable limits can differ from the scalar FP32 clamp path.
Step 3: Optional output quantization (when SFD outputs are generated):
Diagram
API usage
BF16
High-level wrapper
Class API
use_dynamic_sched=False selects static scheduling; True caches a dynamic-M
callable for compatible shapes. Cache keys retain compile-sensitive layouts,
dtypes, activation, dbias policy, scheduler, tiles/clusters, features, and
overlap margin.
Block-scaled
High-level wrapper
Dense mode:
Discrete mode:
Class API
Dense mode:
In the class API, dbias generation is specialized at compile time: if sample_dbias is omitted, dbias_tensor must also be omitted at execute().
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-major or N-major. Must be K-major for FP4.
- 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
- Shape:
-
Input tensor C:
c_tensor/sample_c- Shape:
(valid_m, 2N, 1)— forward activation with interleaved input/gate - Stride:
(2N, 1, valid_m·2N)— must be N-major - Dtype:
{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2}
- Shape:
-
Output tensor D_row:
d_row_tensor/sample_d_row- Shape:
(valid_m, 2N, 1) - Stride:
(2N, 1, valid_m·2N)— must be N-major - Dtype (
d_dtype):{bfloat16, float32}for FP4;{float8_e4m3fn, float8_e5m2}for FP8
- Shape:
-
Output tensor D_col:
d_col_tensor/sample_d_col- Shape:
(valid_m, 2N, 1)— must match D_row dtype and stride
- Shape:
-
Input tensor prob:
prob_tensor/sample_prob- Shape:
(valid_m, 1, 1), dtype:float32
- Shape:
-
Output tensor dprob:
dprob_tensor/sample_dprob- Shape:
(valid_m, 1, 1), dtype:float32 - Must be zero-initialized
- Shape:
-
Scaling tensors:
alphashape(L,),betashape(L,),amaxshape(L, 2, 1),norm_constshape(1,)
Common Parameters
acc_dtype: Must betorch.float32mma_tiler_mn: Kernel tile size. 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 for dSwiGLU and dGeGLU. Default:False. K3-default dSiTU-GLU (situ_beta1=4.0) always uses its packed FP32x2 specialization; non-default dSiTU-GLU uses scalar FP32.m_aligned: Must be256. Default:256discrete_col_sfd: Generate discrete col-major scale factors. Default:Falseact_func: Backward activation function."dswiglu"(default),"dgeglu", or block-scaled"dsituglu"linear_offset: Offset added to the clamped dGeGLU up branch. Default:1.0for dGeGLUgeglu_alpha: dGeGLU sigmoid input scale. Default:1.702glu_clamp_max: dGeGLU upper bound for gate and up. Default:7.0glu_clamp_min: dGeGLU lower bound for up only. Default:-7.0- These activation parameters must match forward and specialize the block-scaled backward cache on Blackwell and Rubin; the BF16 backend retains its fixed alpha/clamp values
situ_beta1: Positive finite gate tanh scale for dSiTU-GLU. Default:4.0situ_beta2: Positive finite up-branch tanh scale for dSiTU-GLU. Default:25.0b_major(discrete only): B tensor major dimension."k"(default) or"n". Must be"k"for FP4.epilogue_op: Optional post-processing.None(default),"identity","relu", or"srelu"
Wrapper-specific Parameters
d_dtype: Output D tensor data type. Default:torch.bfloat16cd_major: Must be"n". Default:"n"n(discrete only): B weight N dimensionb_dtype(discrete only): B weight data type
Wrapper Return Values
Returns a TupleDict (dictionary + tuple unpacking):
d_row_tensor: Row-quantized dGLU outputd_col_tensor: Column-quantized dGLU outputdprob_tensor: Gradient of probamax_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) or K/N-major (discrete). Must be K-major for FP4.C,D_row,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 must have the same dtype
D_rowandD_colmust have the same dtypedbiasmust bebfloat16sf_dtype=float8_e4m3fnis incompatible withsf_vec_size=32- FP8
c_dtypewithvector_f32=Trueis not supported - FP4
ab_dtypeonly supportsd_dtypein{bfloat16, float32} - For non-dbias paths, FP4
ab_dtypewithsf_vec_size=16andd_dtype=float32is not supported - FP8
ab_dtypeonly supportsd_dtypein{float8_e4m3fn, float8_e5m2}
Shapes and Divisibility
Nmust be divisible by 32 (32-column blocks for input/gate interleaving)- Expert count must be
<= 1024 - Each group’s M dimension is aligned to
m_aligned(256) - In the class API,
dbiasis compiled in only whensample_dbiasis provided; passing a runtimedbias_tensorwithoutsample_dbiasraisesValueError 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 backend
- Rubin MXFP8 uses matching FP8 A/B operands, E8M0 scale factors, and
sf_vec_size=32; the dGLU API requires FP8 D outputs with row/column scale factors for these inputs
Usage Examples
For usage examples, see test cases in test/python/fe_api/grouped_gemm/test_grouped_gemm_dglu.py (dense mode, unified API) and test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.py (discrete mode).
Rubin MXFP8 activation-parameter coverage is in
test/python/fe_api/grouped_gemm/test_grouped_gemm_dglu.py
(test_rubin_mxfp8_clamped_dgeglu_*).