Grouped GEMM + GLU (SM100)

View as Markdown

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

Operand contractSelected backend
A and B are BF16BF16
matching supported FP4/FP8 A and B plus scale descriptorsblock-scaled

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, with n=N, b_dtype=torch.bfloat16, and b_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

Cg=αgAgBgT+biasg.C_g = \alpha_g A_g B_g^T + \mathrm{bias}_g.

Columns are paired as alternating 32-wide gate/up blocks. For SwiGLU,

Dg=probg⋅up(Cg)⋅silu(gate(Cg)).D_g = \mathrm{prob}_g \cdot \mathrm{up}(C_g) \cdot \mathrm{silu}(\mathrm{gate}(C_g)).

For GeGLU, let gate = min(gate(C), 7), up = clamp(up(C), -7, 7), geglu_alpha=1.702, and the default linear_offset=1:

Dg=probg⋅(up+linear_offset)⋅gate⋅σ(1.702 gate).D_g = \mathrm{prob}_g \cdot (\mathrm{up}+\mathrm{linear\_offset}) \cdot \mathrm{gate} \cdot \sigma(1.702\,\mathrm{gate}).

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:

  1. Block-scaled grouped GEMM: Low-precision GEMM (FP4, FP8) with per-block scale factors across multiple expert groups
  2. GLU activation: Fused SwiGLU, GeGLU, or SiTU-GLU activation applied to the GEMM output
  3. 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

D=prob [β1tanh⁡(G/β1)σ(G)][β2tanh⁡(U/β2)].D = \mathrm{prob}\, \left[\beta_1\tanh(G/\beta_1)\sigma(G)\right] \left[\beta_2\tanh(U/\beta_2)\right].

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_ptrs shape (num_experts,) of int64
    • 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_ptrs shape (num_experts,) of int64
    • padded_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 (when d_dtype is FP8), shape (32, 4, ceil(valid_m/128), 4, ceil(ceil((N/2)/sf_vec_size)/4), 1)
    • SFD_col: column scale factors (when d_dtype is FP8), shape (32, 4, ceil((N/2)/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1)
    • amax: per-group amax (when d_dtype is 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])):

C[m,n]=αg∑kdequantize(A[m,k],SFA)⋅dequantize(B[n,k,g],SFB)C[m, n] = \alpha_g \sum_{k} \text{dequantize}(A[m, k], \text{SFA}) \cdot \text{dequantize}(B[n, k, g], \text{SFB})

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"):

D[:,bG:(b+1)G]=prob⋅Ub⋅swish(Gb),swish(x)=x⋅σ(x)D[:, bG:(b+1)G] = \text{prob} \cdot U_b \cdot \text{swish}(G_b), \quad \text{swish}(x) = x \cdot \sigma(x)

For GeGLU (act_func="geglu"):

G^b=min⁡(Gb,glu_clamp_max),U^b=clamp⁡(Ub,glu_clamp_min,glu_clamp_max)\widehat{G}_b = \min(G_b, \text{glu\_clamp\_max}), \qquad \widehat{U}_b = \operatorname{clamp}(U_b, \text{glu\_clamp\_min}, \text{glu\_clamp\_max}) D[:,bG:(b+1)G]=prob⋅(U^b+linear_offset)⋅G^b⋅σ(geglu_alpha⋅G^b)D[:, bG:(b+1)G] = \text{prob} \cdot (\widehat{U}_b + \text{linear\_offset}) \cdot \widehat{G}_b \cdot \sigma(\text{geglu\_alpha} \cdot \widehat{G}_b)

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):

SFD_row[m,n]=norm_const⋅max⁡k∈block∣D[m,k]∣⋅rcp_max\text{SFD\_row}[m, n] = \text{norm\_const} \cdot \max_{k \in \text{block}} |D[m, k]| \cdot \text{rcp\_max} Dquantized[m,n]=D[m,n]⋅norm_constSFD_row[m,n]D_{\text{quantized}}[m, n] = D[m, n] \cdot \frac{\text{norm\_const}}{\text{SFD\_row}[m, n]}

Diagram

A (valid_m×K×1) B (N×K×L) or b_ptrs padded_offsets
SFA SFB or sfb_ptrs |
| | |
| +------------+ |
| | |
v v v
Dequantize → Grouped GEMM (per group ranges) → Select B[:,:,group_idx]
|
| × alpha[group_idx]
v
C (valid_m×N×1)
|
| Pair 32-col blocks: [G0|U0|G1|U1|...]
| U_b × swish(G_b) [SwiGLU]
| clamp gate/up, then
| (U_b+offset) × G_b·σ(alpha·G_b) [GeGLU]
v
| × prob
v
D (valid_m×N/2×1)
|
+----------+-----------+
| |
v v
Row Quantize Col Quantize
| |
v v
D_row, SFD_row D_col, SFD_col

API usage

BF16

High-level wrapper

import cudnn
import torch
# Dense wrapper. The required scale positions are explicitly None for BF16.
out = cudnn.grouped_gemm_glu_wrapper_sm100(
a_tensor=a,
sfa_tensor=None,
padded_offsets=padded_offsets,
alpha_tensor=alpha,
b_tensor=b,
sfb_tensor=None,
bias_tensor=bias,
prob_tensor=prob,
act_func="swiglu",
generate_c=True,
use_dynamic_sched=True,
)
c, d, d_col, amax, sfd_row, sfd_col = out
# Discrete wrapper.
out = cudnn.grouped_gemm_glu_wrapper_sm100(
a_tensor=a,
sfa_tensor=None,
padded_offsets=padded_offsets,
alpha_tensor=alpha,
b_ptrs=b_ptrs,
sfb_ptrs=None,
n=N,
b_dtype=torch.bfloat16,
prob_tensor=prob,
act_func="geglu",
)

Class API

op = cudnn.GroupedGemmGluSm100(
sample_a=a,
sample_c=c,
sample_d=d,
sample_d_col=None,
sample_sfa=None,
sample_padded_offsets=padded_offsets,
sample_alpha=alpha,
sample_b=b,
sample_sfb=None,
sample_prob=prob,
act_func="swiglu",
generate_c=True,
)
assert op.check_support()
op.compile()
op.execute(
a_tensor=a, c_tensor=c, d_tensor=d, sfa_tensor=None,
padded_offsets=padded_offsets, alpha_tensor=alpha,
b_tensor=b, sfb_tensor=None, prob_tensor=prob,
)

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:

from cudnn import grouped_gemm_glu_wrapper_sm100
from cuda.bindings import driver as cuda
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
outputs = grouped_gemm_glu_wrapper_sm100(
a_tensor=a,
sfa_tensor=sfa,
padded_offsets=padded_offsets,
alpha_tensor=alpha,
bias_tensor=bias,
# Dense mode weights:
b_tensor=b,
sfb_tensor=sfb,
# Common:
norm_const_tensor=norm_const,
prob_tensor=prob,
acc_dtype=torch.float32,
c_dtype=torch.bfloat16,
d_dtype=torch.bfloat16,
cd_major="n",
mma_tiler_mn=(256, 256),
cluster_shape_mn=(2, 1),
sf_vec_size=32,
vector_f32=False,
m_aligned=256,
discrete_col_sfd=False,
act_func="swiglu",
current_stream=stream,
)
# dictionary access:
c = outputs["c_tensor"]
d = outputs["d_tensor"]
d_col = outputs["d_col_tensor"]
amax = outputs["amax_tensor"]
sfd_row = outputs["sfd_row_tensor"]
sfd_col = outputs["sfd_col_tensor"]
# or tuple unpacking:
c, d, d_col, amax, sfd_row, sfd_col = outputs

Discrete mode:

outputs = grouped_gemm_glu_wrapper_sm100(
a_tensor=a,
sfa_tensor=sfa,
padded_offsets=padded_offsets,
alpha_tensor=alpha,
# Discrete mode weights:
b_ptrs=b_ptrs, # int64 tensor of per-expert B data pointers
sfb_ptrs=sfb_ptrs, # int64 tensor of per-expert SFB data pointers
n=n_dim, # B weight N dimension
b_dtype=torch.uint8, # B weight data type
b_major="k", # B tensor major dimension
# Common:
norm_const_tensor=norm_const,
prob_tensor=prob,
act_func="geglu", # GeGLU activation
current_stream=stream,
)

bias_tensor must use the kernel layout expected by the fused bias path: shape (N, L) and stride (1, N).

Class API

Dense mode:

from cudnn import GroupedGemmGluSm100
from cuda.bindings import driver as cuda
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
api = GroupedGemmGluSm100(
sample_a=a,
sample_c=c,
sample_d=d,
sample_sfa=sfa,
sample_padded_offsets=padded_offsets,
sample_alpha=alpha,
sample_d_col=d_col,
sample_bias=bias,
# Dense mode:
sample_b=b,
sample_sfb=sfb,
# Optional quantization outputs
sample_sfd_row=sfd_row,
sample_sfd_col=sfd_col,
sample_amax=amax,
sample_norm_const=norm_const,
sample_prob=prob,
# Configuration
acc_dtype=torch.float32,
mma_tiler_mn=(256, 256),
cluster_shape_mn=(2, 1),
sf_vec_size=32,
act_func="swiglu",
)
assert api.check_support()
api.compile()
api.execute(
a_tensor=a, c_tensor=c, d_tensor=d,
sfa_tensor=sfa, padded_offsets=padded_offsets, alpha_tensor=alpha,
b_tensor=b, sfb_tensor=sfb, bias_tensor=bias,
d_col_tensor=d_col, sfd_row_tensor=sfd_row, sfd_col_tensor=sfd_col,
amax_tensor=amax, norm_const_tensor=norm_const, prob_tensor=prob,
current_stream=stream,
)

sample_bias and runtime bias_tensor must both use shape (N, L) and stride (1, N).

Discrete mode:

api = GroupedGemmGluSm100(
sample_a=a,
sample_c=c,
sample_d=d,
sample_sfa=sfa,
sample_padded_offsets=padded_offsets,
sample_alpha=alpha,
sample_d_col=d_col,
# Discrete mode:
num_experts=num_experts,
b_shape=(n, k),
b_dtype=torch.uint8,
# Configuration
act_func="geglu",
b_major="k",
)
assert api.check_support()
api.compile()
api.execute(
a_tensor=a, c_tensor=c, d_tensor=d,
sfa_tensor=sfa, padded_offsets=padded_offsets, alpha_tensor=alpha,
b_ptrs=b_ptrs, sfb_ptrs=sfb_ptrs,
d_col_tensor=d_col, prob_tensor=prob,
current_stream=stream,
)

Parameters

Weight Mode

The weight mode is auto-detected from constructor arguments:

  • Dense: Provide sample_b and sample_sfb (contiguous weight tensors)
  • Discrete: Provide num_experts, b_shape, and b_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}
  • Input tensor B (dense mode): b_tensor / sample_b

    • Shape: (N, K, L) where L = num_groups
    • Stride: (K, 1, N·K) — must be K-major
    • Dtype: Must match A
  • 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")
  • Output tensor C: returned in wrapper dict or c_tensor in 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
  • 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
  • Output tensor D_col: d_col_tensor / sample_d_col

    • Shape: (valid_m, N/2, 1) — must match D dtype and stride
  • 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)
  • Group offsets: padded_offsets — shape (L,), dtype int32

  • Scaling tensors: alpha shape (L,), prob shape (valid_m, 1, 1), amax shape (L, 1), norm_const shape (1,)

Common Parameters

  • acc_dtype: Must be torch.float32
  • mma_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) when TILE_M=256, (1, 1) otherwise
  • sf_vec_size: Scale factor vector size. {16, 32}. Default: 16
  • vector_f32: Enable packed f32 operations. Default: False
  • m_aligned: Must be 256 (FIX_PAD_SIZE). Default: 256
  • discrete_col_sfd: Generate discrete col-major scale factors. Default: False
  • act_func: Activation function. "swiglu" (default), "geglu", or block-scaled "situglu"
  • linear_offset: Offset added to the clamped GeGLU up branch. Default: 1.0 for GeGLU
  • geglu_alpha: GeGLU sigmoid input scale. Default: 1.702
  • glu_clamp_max: GeGLU upper bound for gate and up. Default: 7.0
  • glu_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.0
  • situ_beta2: Positive finite up-branch tanh scale for SiTU-GLU. Default: 25.0
  • b_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.bfloat16
  • d_dtype: Output D tensor data type. Default: torch.bfloat16
  • cd_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 result
  • d_tensor: Row-quantized GLU output
  • d_col_tensor: Column-quantized GLU output
  • amax_tensor: Per-group amax (when d_dtype is 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

  • A must be K-major
  • B must be K-major (dense mode). For discrete mode: K-major or N-major (K-major required for FP4)
  • C, D, D_col must be N-major
  • All tensors must be 16-byte aligned along the contiguous dimension

Data Types

Formatab_dtypesf_dtypesf_vec_sized_dtype
MXFP8float8_e4m3fn or float8_e5m2float8_e8m0fnu32{float16, bfloat16, float8_e4m3fn, float8_e5m2}
NVF4float4_e2m1fn_x2 or uint8{float8_e4m3fn, float8_e8m0fnu}{16, 32}{bfloat16, float32}

Additional Type Constraints

  • A and B must have the same dtype
  • Scale factor tensors (SFA, SFB, SFD_row, SFD_col) must have the same dtype
  • D and D_col must have the same dtype
  • bias must be one of {float16, bfloat16, float32}
  • bias must have shape (N, L) and stride (1, N)
  • For non-bias paths, FP4 ab_dtype with sf_vec_size=16 and d_dtype=float32 is not supported
  • FP4 ab_dtype requires c_dtype in {float16, bfloat16}

Shapes and Divisibility

  • N must 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=True is supported only by the block-scaled kernel with exactly one expert. In this mode the kernel derives padded_offsets[0] from runtime A.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).