Grouped GEMM + dSwiGLU (SM100)
Grouped GEMM + dSwiGLU (SM100)
This is an experimental API and subject to change.
Legacy contiguous-only API note: This page documents the older contiguous-only dSwiGLU API. For new integrations, prefer the unified Grouped GEMM + dGLU API, which covers dense and discrete weight layouts.
JAX support
cudnn.grouped_gemm_dswiglu_wrapper_sm100 accepts Torch tensors and canonical
MXFP8 JAX arrays or tracers. Torch execution is unchanged; JAX dispatches to
cudnn.jax.grouped_gemm_dswiglu, eagerly or under jax.jit, with XLA-owned
buffers and stream ordering. cudnn.torch.grouped_gemm_dswiglu remains an
alias to the same wrapper and inherits its dispatch. Direct API-class construction
with JAX samples remains unsupported.
Wrapper signatures and defaults are unchanged. JAX wrapper calls must set
sf_vec_size=32 and an explicit FP8 d_dtype. The direct cudnn.jax API retains
its MXFP8 defaults. See the JAX execution contract below.
Overview
Grouped GEMM + dSwiGLU fusion: A contiguous grouped block-scaled GEMM fused with a dSwiGLU backward epilogue on NVIDIA Blackwell GPUs (SM100+), designed for MoE (Mixture of Experts) workloads. Implemented with CUTLASS/CUTE.
Groups are contiguous in the M dimension and described by padded_offsets (cumulative aligned end offsets).
This kernel performs:
- Block-scaled grouped GEMM: Low-precision GEMM (FP4, FP8) with per-block scale factors across multiple expert groups
- dSwiGLU 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
- Inputs
A: contiguous activation tensor across all groups, shape(valid_m, K, 1)B: weight tensor across all groups, shape(N, K, L)C: 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: scale factor tensor for B, shape(32, 4, ceil(N/128), 4, ceil(ceil(K/sf_vec_size)/4), L)padded_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 forC, shape(L,)prob: per-row gating probabilities, shape(valid_m, 1, 1)norm_const: normalization constant for FP8 quantization, shape(1,)
- Outputs
D_row: row-quantized dSwiGLU output, shape(valid_m, 2N, 1)D_col: column-quantized dSwiGLU output, shape(valid_m, 2N, 1)dprob: gradient ofprob, shape(valid_m, 1, 1). Must be zero-initialized.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: dSwiGLU backward epilogue (performed with 32-column interleaving along 2N):
- Scale
Cbybeta_gper group and deinterleave into input/gate halves by 32-wide blocks. swish = gate * sigmoid(gate)dprobis the sum over 32-column chunks ofswish * input * refab = ref * prob * swishdswiglu = ref * prob * input * sigmoid(gate) * (1 + gate * (1 - sigmoid(gate)))- Interleave
[ab, dswiglu]back intoD_row/D_colwith 32-column blocks.
Step 3: Optional output quantization (when SFD outputs are generated):
Diagram
API Usage
High-level Wrapper
Class API
Parameters
Input/Output Tensors
-
Input tensor A:
a_tensor(wrapper) orsample_a,a_tensor(class)- 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}uint8is interpreted as packed FP4 (two FP4 values per byte)
- Shape:
-
Input tensor B:
b_tensor(wrapper) orsample_b,b_tensor(class)- Shape:
(N, K, L)whereL = num_groups - Stride:
(K, 1, N·K)(K-major) or(1, N, N·K)(N-major). Must be K-major for fp4 inputs. - Dtype (
ab_dtype): Must match A
- Shape:
-
Input tensor C:
c_tensor(wrapper) orsample_c,c_tensor(class)- Shape:
(valid_m, 2N, 1) - Stride:
(2N, 1, valid_m·2N)- must be N-major - Dtype (
c_dtype):{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2}
- Shape:
-
Output tensor D_row:
d_row_tensor(class) or returned in wrapper dict- Shape:
(valid_m, 2N, 1) - Stride:
(2N, 1, valid_m·2N)- must be N-major - Dtype (
d_dtype):{bfloat16, float32}for FP4 inputs;{float8_e4m3fn, float8_e5m2}for FP8 inputs
- Shape:
-
Output tensor D_col:
d_col_tensor(class) or returned in wrapper dict- Shape:
(valid_m, 2N, 1) - Stride:
(2N, 1, valid_m·2N)- must match D_row (N-major) - Dtype: Must match D_row
- Shape:
-
Input tensor prob:
prob_tensor(wrapper) orsample_prob(class)- Shape:
(valid_m, 1, 1) - Dtype:
float32
- Shape:
-
Output tensor dprob:
dprob_tensor(wrapper) orsample_dprob(class)- Shape:
(valid_m, 1, 1) - Dtype:
float32 - Must be zero-initialized.
- Shape:
-
Scale factor tensors
- SFA (A scale factor):
sfa_tensor(wrapper) orsample_sfa,sfa_tensor(class)- Shape:
(32, 4, ceil(valid_m/128), 4, ceil(ceil(K/sf_vec_size)/4), 1) - Dtype (
sf_dtype):{float8_e8m0fnu, float8_e4m3fn}
- Shape:
- SFB (B scale factor):
sfb_tensor(wrapper) orsample_sfb,sfb_tensor(class)- Shape:
(32, 4, ceil(N/128), 4, ceil(ceil(K/sf_vec_size)/4), L) - Dtype: Must match SFA
- Shape:
- SFD_row (D row scale factor, optional):
sfd_row_tensor(wrapper) orsample_sfd_row,sfd_row_tensor(class)- Shape:
(32, 4, ceil(valid_m/128), 4, ceil(ceil((2N)/sf_vec_size)/4), 1) - Dtype: Must match SFA
- Required when: SFD outputs are enabled (FP8 inputs)
- Shape:
- SFD_col (D column scale factor, optional):
sfd_col_tensor(wrapper) orsample_sfd_col,sfd_col_tensor(class)- Shape:
(32, 4, ceil((2N)/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1) - Dtype: Must match SFA
- Required when: SFD outputs are enabled (FP8 inputs)
- Shape:
- SFA (A scale factor):
-
Group offsets
- padded_offsets: Cumulative sum of aligned group M sizes
- Shape:
(L,)whereL = num_groups - Dtype:
int32 padded_offsets[-1]equalsvalid_m; each offset is a multiple ofm_aligned
- Shape:
- padded_offsets: Cumulative sum of aligned group M sizes
-
Scaling tensors
- alpha: Per-group scaling factors
- Shape:
(L,)whereL = num_groups - Dtype:
float32
- Shape:
- beta: Per-group scaling factors for
C- Shape:
(L,)whereL = num_groups - Dtype:
float32
- Shape:
- amax (optional): Per-group max absolute values
- Shape:
(L, 2, 1) - Dtype:
float32 - Required when:
d_dtype ∈ {bfloat16, float16}
- Shape:
- norm_const (optional): Normalization constant for FP8 quantization
- Shape:
(1,) - Dtype:
float32 - Required when:
sfd_row_tensor/sfd_col_tensorare provided (FP8 inputs)
- Shape:
- alpha: Per-group scaling factors
Common Parameters
-
acc_dtype: torch.dtype- Accumulator dtype. Must be
torch.float32
- Accumulator dtype. Must be
-
mma_tiler_mn: Tuple[int, int]- Kernel tile size
(TILE_M, TILE_N). Default:(256, 256) TILE_M ∈ {64, 128, 256}TILE_N ∈ {128, 256}
- Kernel tile size
-
cluster_shape_mn: Tuple[int, int] | None- Thread Block cluster shape
(CLUSTER_M, CLUSTER_N) - Constraints: positive powers of 2, both <= 4,
CLUSTER_M × CLUSTER_N <= 16 - Default:
(2, 1)whenTILE_M=256,(1, 1)otherwise
- Thread Block cluster shape
-
sf_vec_size: int- Scale factor vector size (number of elements per scale factor)
- Allowed values:
{16, 32}. Default:16
-
vector_f32: bool- Enable packed f32 operations for improved performance
- Default:
False
-
m_aligned: int- Alignment requirement for group M dimension
- Must equal
FIX_PAD_SIZE(256) and be divisible bymma_tiler_mn[0] - Default:
256
-
discrete_col_sfd: bool- If True, generate discrete column scale factors grouped by expert tiles
- Only applies when
sfd_row_tensor,sfd_col_tensor, andnorm_const_tensorare provided - No extra inputs are required; this only changes the layout of
sfd_col_tensor - Default:
False
-
epilogue_op: Optional[str]- Optional epilogue operation. Valid values:
None,"none","identity","relu","srelu" - Default:
None
- Optional epilogue operation. Valid values:
-
CUDA stream (
current_streamin class API,current_streamin wrapper)
Wrapper-specific Parameters: grouped_gemm_dswiglu_wrapper_sm100
d_dtype: torch.dtype: Output D tensor data type. Default:torch.bfloat16cd_major: str: Major dimension for C and D tensors. Must be"n"(only N-major layout is supported). Default:"n"
Wrapper Return Values
Returns a TupleDict - a dictionary-like object that also supports tuple unpacking and integer indexing.
Dictionary keys (also the tuple unpacking order):
d_row_tensor: Row-quantized dSwiGLU outputd_col_tensor: Column-quantized dSwiGLU outputdprob_tensor: Gradient ofprobamax_tensor: Per-group amax (whend_dtype ∈ {bfloat16, float16})sfd_row_tensor: Row scale factors (when SFD outputs are enabled)sfd_col_tensor: Column scale factors (when SFD outputs are enabled)
Class-specific Parameters
GroupedGemmDswigluSm100 (constructor)
sample_a,sample_b,sample_c,sample_d_row,sample_d_col,sample_sfa,sample_sfb,sample_padded_offsets,sample_alpha,sample_beta,sample_prob,sample_dprob,sample_sfd_row,sample_sfd_col,sample_amax,sample_norm_const- see Input/Output tensors- Note:
sample_sfd_row,sample_sfd_col,sample_norm_constmust be allNoneor all notNone
- Note:
GroupedGemmDswigluSm100.execute
a_tensor,b_tensor,c_tensor,d_row_tensor,d_col_tensor,sfa_tensor,sfb_tensor,padded_offsets,alpha_tensor,beta_tensor,prob_tensor,dprob_tensor,sfd_row_tensor,sfd_col_tensor,amax_tensor,norm_const_tensor- see Input/Output tensors. Must have same layout as sample tensors provided in constructor.
Support Surface and Constraints
Layouts and Strides
Amust be K-major (contiguous along K dimension)Bmust be K-major (contiguous along K dimension) or N-major (contiguous along N dimension). Must be K-major for fp4 inputs.C,D_row, andD_colmust be N-major (contiguous along N dimension)- All tensors must be 16-byte aligned along the contiguous dimension
Canonical layouts (additive)
Each input is also accepted in its natural row-major form. Canonical inputs compile at their own rank and bind directly, with no per-call host-side views; the pre-permuted kernel-facing forms above keep working unchanged:
A:(valid_m, K)row-majorB:(L, N, K)C-contiguousC:(valid_m, 2N)row-majorSFA/SFB: any dense C-contiguous buffer with the MMA-tiled element count, e.g. flat 1-D or the physical(L, ceil(mn/128), ceil(ceil(K/sf_vec_size)/4), 32, 4, 4)allocation. The kernel rebuilds the MMA-tiled SF layouts from the GEMM shapes and reads only the base pointer.prob:(valid_m,),float32orbfloat16alpha_tensorremains required; pass explicit per-group scaling factors.
Flat SF buffers must already contain the packed MMA-tiled scale bytes in physical order. Ordinary row-major logical scales need packing before this API is called.
These layouts are also used by the JAX execution path below. Unified GLU/dGLU APIs are separate.
When A is canonical (2-D), the wrapper returns natural-shaped outputs:
d_row/d_col (valid_m, 2N) row-major, dprob (valid_m,), and
sfd_row/sfd_col as C-contiguous physical (1, ceil(mn/128), rest, 32, 4, 4) buffers.
JAX execution
cudnn.jax.grouped_gemm_dswiglu runs the contiguous-weight MXFP8 fusion
through cudnn.jax.call, eagerly or under jax.jit. All operands are ordinary
JAX arrays managed by XLA. Torch callers use cudnn.torch.grouped_gemm_dswiglu
or the existing top-level wrapper name.
Both paths return TupleDict with the same key order and tuple-unpacking behavior.
The JAX path registers this output type as a JAX pytree.
Use canonical A (m,k), B (experts,n,k), and prob (m,) (fp32 or bf16).
Scale factors are E8M0 arrays, or uint8 bit patterns, containing the packed
MMA-tiled physical bytes; physical 6-D and flat buffers are accepted. Pass explicit
fp32 alpha (experts,), norm_const (1,), and int32 padded_offsets (experts,).
Offsets must be nondecreasing multiples of 256 in [0,m]; m must be a positive
multiple of 256. These device values are the caller’s responsibility.
Backward also requires saved C (m,2n) and explicit fp32 beta (experts,).
It returns d_row_tensor, d_col_tensor, dprob_tensor, physical
sfd_row_tensor/sfd_col_tensor, and amax_tensor=None.
The JAX API fixes scale-vector size to 32 and defaults d_dtype to FP8 e4m3.
It requires explicit probability and normalization arrays; backward also requires
beta. Mixed Torch/JAX operands are rejected. Torch-specific streams, output buffers,
accumulation/layout options, and epilogues are not JAX parameters. The optional JAX
configuration is d_dtype, mma_tiler_mn, and cluster_shape_mn.
Configuration arguments must be static under jax.jit:
The existing wrapper also works under jax.jit:
On the wrapper’s JAX path, unsupported options raise ValueError: non-FP32
accumulation, non-n output layout, scale-vector size other than 32,
vector_f32=True, non-default m_aligned, discrete_col_sfd=True, and caller
streams, plus caller output buffers and non-identity backward epilogues.
The Torch alias preserves the existing wrapper signature, including its dtype and scale-vector defaults. For example, select MXFP8 explicitly:
This initial bridge supports FP8 e4m3/e5m2 A/B and e4m3 D, with E8M0 block
scales of vector size 32. The packed backward quantizer does not support e5m2 D.
Packed FP4, BF16 D, bias, and discrete-column SF layout
are outside its contract. Outputs are initialized to zero (raw zero bytes for SF)
to define untouched padding; backward dprob also requires initialization for atomic
accumulation. CUDA graph compatibility uses the standard CuTeDSL JAX bridge.
This API supplies the fused backward operation explicitly; it does not register
an automatic jax.grad rule. Full TE training integration is separate validation.
Data Types
Input/Weight Types (ab_dtype)
Additional Type Constraints
AandBmust have the same dtypeSFA,SFB,SFD_row, andSFD_colmust have the same dtypeD_rowandD_colmust have the same dtypeacc_dtypemust befloat32sf_dtype=float8_e4m3fnis incompatible withsf_vec_size=32- FP8
c_dtypewithvector_f32=Trueis not supported - FP4
ab_dtypeonly supportsd_dtype ∈ {bfloat16, float32} - FP8
ab_dtypeonly supportsd_dtype ∈ {float8_e4m3fn, float8_e5m2}
Scale Factor Output Requirements
-
When
sfd_row_tensor/sfd_col_tensorare provided (FP8 inputs):sfd_row_tensor,sfd_col_tensor, andnorm_const_tensorare all required- These must be provided together (all None or all not None)
-
When
d_dtype ∈ {bfloat16, float16}:amax_tensoris required for tracking per-group max values
Tiling and Cluster
mma_tiler_mn[0] = 256enables 2-CTA instructions automatically (use_2cta_instrs=True)- When
use_2cta_instrs=True:cluster_shape_mn[0]must be divisible by 2 m_alignedmust be divisible bymma_tiler_mn[0]to prevent tiles from spanning multiple groups
Shapes and Divisibility
Nmust be divisible by 32 (32-column blocks for input/gate interleaving)padded_offsetslengthLis the expert count and must be<= 1024- Each group’s M dimension is aligned to
m_aligned valid_m = padded_offsets[-1]determines the actual tensor M dimension- Scale factor tensor shapes follow the MMA atom tiling pattern:
(32, 4, ceil(dim/128), 4, ceil(K_groups/4), L)
Environment
- Requires CUDA with SM100+ compute capability (Blackwell GPUs)
Usage Examples
For usage examples, see test cases in test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu.py + test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_utils.py