GEMM + SwiGLU (SM100)
GEMM + SwiGLU (SM100)
This is an experimental API and subject to change.
Overview
GEMM + SwiGLU fusion: A persistent, batched dense GEMM fused with a SwiGLU epilogue on NVIDIA Blackwell GPUs (SM100+), implemented with CUTLASS/CUTE. It produces both the full GEMM output AB12 and a SwiGLU-projected tensor C in a single pass.
This API supports two modes:
- Standard mode: High-precision GEMM with SwiGLU epilogue
- Quantized mode (block-scaled): Low-precision GEMM using block scaling supporting FP4 and FP8 data types
Shapes
- Inputs:
A: shape(M, K, L)B: shape(N, K, L)
- Outputs:
-
AB12: shape(M, N, L)– full GEMM result -
C: shape(M, N/2, L)– SwiGLU-projected resultLis the batch dimension.
-
Equations
- GEMM (per batch l):
-
SwiGLU epilogue (performed by pairing 32-column blocks along
N):Let block size
G = 32. For each pair of consecutive 32-wide column blocks inAB12:- Input block:
X_b = AB12[:, 2*b*G : 2*b*G + G, :] - Gate block:
G_b = AB12[:, 2*b*G + G : 2*b*G + 2*G, :]
- Input block:
Notes:
- The
alphascaling is applied before the SwiGLU; bothX_bandG_bare from the scaled GEMM results. AB12stores the entire scaled GEMM output (both input and gate blocks), whileCstores the fused SwiGLU-projected result with half the columns.- N divisibility requirement:
Nmust be divisible by 64 (two consecutive 32-column blocks) to ensure proper pairing for the SwiGLU operation.
Diagram
API Usage
The tensor parameters are type-erased: torch tensors and JAX arrays are both accepted (torch is only imported when torch tensors/dtypes are passed, jax only when JAX arrays are passed). Dtype parameters accept torch dtypes, numpy/ml_dtypes dtypes, dtype name strings, or cutlass types. The JAX contract matches gemm_amax (see gemm_amax.md “Using JAX arrays”): A/B k-major (M, K, 1)/(N, K, 1), outputs n-major only, batch L == 1, SF tensors accepted in the physical C-contiguous atom shape (L, MN', K', 32, 4, 4); the eager entry points run on the CUDA legacy default stream (synchronize before reading outputs).
For jitted JAX programs, use gemm_swiglu_jax_sm100 — an XLA custom call (built on cudnn.jax.call / CuTeDSL’s native cutlass.jax bridge, jax dependency group) that runs on XLA’s compute stream, returns fresh (ab12, c) arrays, and composes with jax.jit. Both the standard and the blockscaled MXFP8 quantized kernels are supported. alpha is a static (trace-time) parameter.
High-level wrapper (Standard Mode)
High-level wrapper (Quantized Mode)
When scale factor tensors are provided, the wrapper uses the block-scaled quantized kernel.
Class API (Standard Mode)
Class API (Quantized Mode)
Parameters
Input/Output tensors
- Input tensor A:
a_tensor(wrapper) orsample_a,a_tensor(class)- Shape:
(M, K, L) - Stride:
(1, M, M·K)form-major or(K, 1, M·K)fork-major- Quantized mode: Must be
k-major for FP4 inputs
- Quantized mode: Must be
- Dtype (
ab_dtype):- Standard mode:
{float16, bfloat16, float32, float8_e4m3fn, float8_e5m2} - Quantized mode:
{float4_e2m1fn_x2, uint8, float8_e4m3fn, float8_e5m2}uint8is interpreted as packed FP4 (two FP4 values per byte)
- Standard mode:
- Shape:
- Input tensor B:
b_tensor(wrapper) orsample_b,b_tensor(class)- Shape:
(N, K, L) - Stride:
(1, N, N·K)forn-major or(K, 1, N·K)fork-major - Dtype (
ab_dtype): Must matchA
- Shape:
- Output tensor AB12:
result["ab12_tensor"](wrapper) orsample_ab12,ab12_tensor(class)- Shape:
(M, N, L) - Stride:
(1, M, M·N)form-major or(N, 1, M·N)forn-major. Provided asc_majorargument for wrapper- Quantized mode: Must be
n-major for FP4 outputs
- Quantized mode: Must be
- Dtype (
ab12_dtype, provided asab12_dtypeargument for wrapper):- Standard mode:
{float32, float16, bfloat16}ifacc_dtype == float32,{float16, bfloat16}ifacc_dtype == float16 - Quantized mode:
{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2}
- Standard mode:
- Shape:
- Output tensor C:
result["c_tensor"](wrapper) orsample_c,c_tensor(class)- Shape:
(M, N/2, L) - Stride:
(1, M, M·N/2)form-major or(N/2, 1, M·N/2)forn-major. Must match withAB12 - Dtype (
c_dtype, provided asc_dtypeargument for wrapper):- Standard mode:
{float16, bfloat16} - Quantized mode:
{float32, float16, bfloat16, float8_e4m3fn, float8_e5m2}
- Standard mode:
- Shape:
- Quantization-specific tensors
- Input tensor SFA (A scale factor):
sfa_tensor(wrapper) orsample_sfa,sfa_tensor(class)- Shape:
(32, 4, ceil(M/128), 4, ceil(ceil(K/sf_vec_size)/4), L) - Dtype:
{float8_e8m0fnu, float8_e4m3fn}
- Shape:
- Input tensor 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:
- Output tensor SFC (C scale factor, Optional):
result["sfc_tensor"](wrapper) orsample_sfc,sfc_tensor(class)- Shape:
(32, 4, ceil(M/128), 4, ceil(ceil((N/2)/sf_vec_size)/4), L) - Dtype: Must match
SFA - Required when:
c_dtype ∈ {float8_e4m3fn, float8_e5m2}
- Shape:
- Output tensor AMAX (Optional):
result["amax_tensor"](wrapper) orsample_amax,amax_tensor(class)- Shape:
(1,) - Dtype:
float32 - Required when:
ab_dtypeis FP4 andc_dtype == bfloat16
- Shape:
- Input tensor Norm Const (Optional):
norm_const_tensor(wrapper) orsample_norm_const,norm_const_tensor(class)- Shape:
(1,) - Dtype:
float32 - Required when:
c_dtype ∈ {float8_e4m3fn, float8_e5m2}
- Shape:
- Input tensor SFA (A scale factor):
Common parameters
alpha: float- Scalar multiplier applied to the GEMM result before SwiGLU.
- Default:
1.0
acc_dtype: torch.dtype- Accumulator dtype.
- Standard mode:
{float32, float16}. Default:torch.float32 - Quantized mode: Must be
float32
mma_tiler_mn: Tuple[int, int]- Kernel tile size
(TILE_M, TILE_N). Default:(128, 128) TILE_M ∈ {128, 256}- Standard mode:
TILE_N ∈ {32, 64, ..., 224, 256} - Quantized mode:
TILE_N ∈ {64, 128, 192, 256}
- Kernel tile size
cluster_shape_mn: Tuple[int, int] | None- Thread Block cluster shape
(CLUSTER_M, CLUSTER_N) - Constraints: positive powers of 2,
CLUSTER_M*CLUSTER_N ≤ 16. - Default:
(1,1)ifmma_tiler_mn[0] != 256else(2,2).
- Thread Block cluster shape
- CUDA stream (
current_streamin class API,streamin wrapper) - Quantization-specific parameters
sf_vec_size: int- Scale factor vector size (number of elements per scale factor)
- Allowed values:
{16, 32}. Default:16 - Constraints:
- FP8 inputs require
sf_vec_size=32withsf_dtype=float8_e8m0fnu - FP4 inputs do not support
sf_vec_size=32withsf_dtype=float8_e4m3fn
- FP8 inputs require
vector_f32: bool- Enable packed f32 operations for improved performance
- Default:
False
ab12_stages: int- Number of pipeline stages for AB12 output
- Default:
4
Wrapper-specific parameters: gemm_swiglu_wrapper_sm100
a_tensor,b_tensor: see Input/Output tensorsc_major: str: see Input/Output tensors. Default:"n"ab12_dtype: torch.dtype: see Input/Output tensors. Default:torch.float32c_dtype: torch.dtype: see Input/Output tensors. Default:torch.float16sfa_tensor,sfb_tensor,norm_const_tensor: see Quantization-specific tensorssf_vec_size,vector_f32,ab12_stages: see Quantization-specific parameters
Wrapper return values
Returns a TupleDict with fixed keys:
ab12_tensor: Intermediate GEMM outputc_tensor: SwiGLU outputsfc_tensor: Output scale factors (orNonewhen not applicable)amax_tensor: Max-abs output (orNonewhen not applicable)
Tuple unpacking order is always:
(ab12_tensor, c_tensor, sfc_tensor, amax_tensor).
- Standard mode:
sfc_tensor is Noneandamax_tensor is None - Quantized mode:
sfc_tensorand/oramax_tensorare populated based on dtype/configuration
Class-specific parameters
GemmSwigluSm100 (constructor)
sample_a,sample_b,sample_ab12,sample_c– see Input/Output tensorssample_sfa,sample_sfb,sample_sfc,sample_amax,sample_norm_const– see Scale factor tensors (quantized mode)
GemmSwigluSm100.execute
a_tensor,b_tensor,ab12_tensor,c_tensor– see Input/Output tensors. Must have same layout as sample tensors provided in constructor.sfa_tensor,sfb_tensor,sfc_tensor,amax_tensor,norm_const_tensor– see Scale factor tensors (quantized mode)
Support surface and constraints
Layouts and strides
AB12andCmust have the same major order.A,B,AB12must be 16-byte aligned along the contiguous dimension.- For FP4 inputs (quantized mode):
AandBmust bek-major,AB12must ben-major.
Dtypes
Standard mode
A/Bmust have the same dtype.ab12_dtype ∈ {float8_e4m3fn, float8_e5m2}is currently disabledacc_dtype == float16is only supported withab_dtype ∈ {float16, float8_e4m3fn, float8_e5m2}ab12_dtype ∈ {float32}requiresacc_dtype == float32
Quantized mode
The quantized kernel supports the following configurations:
Additional constraints:
acc_dtypemust befloat32- Not compatible with FP8 c_dtype. BF16
c_dtypeis expected. - For MXFP8 inputs, ab12_dtype` should be float16 or bfloat16.
- When
c_dtype ∈ {float8_e4m3fn, float8_e5m2}:sfc_tensorandnorm_const_tensorare required - When
ab_dtypeis FP4 andc_dtype == bfloat16:amax_tensoris required c_dtypeandab12_dtypecannot both befloat32
Tiling and cluster
- Using
TILE_M == 256requiresmma_tiler_mn[0] == 256(enables 2-CTA instructions). - If
TILE_M == 128andcluster_shape_mn != (1, 1),mma_tiler_mnmust be exactly(128, 128). - If
mma_tiler_mn[0] == 256,CLUSTER_Mmust be divisible by 2 - Standard mode: If
mma_tiler_mn[0] != 256,cluster_shape_mnmust be(1, 1).
Environment
- Requires CUDA with SM100+ compute capability
Usage examples
For usage examples, see test cases in test/python/fe_api/gemm/test_gemm_swiglu.py