Discrete Grouped GEMM + dSwiGLU (SM100)
Discrete Grouped GEMM + dSwiGLU (SM100)
This is an experimental API and subject to change.
JAX support
Supports JAX arrays in FP8 configurations: b_ptrs/sfb_ptrs as packed-uint8 (or x64 int64) pointer arrays, SFA in the physical C-contiguous atom shape, SFD outputs allocated the same way (the kernel rebuilds all SF layouts from the GEMM shapes and reads only base pointers). Packed-fp4 inputs 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 discrete_grouped_gemm_dswiglu_jax_sm100 (built on cudnn.jax.call; k-major weights only): all outputs (d_row/d_col, SFD tensors, amax, dprob as a bridge-managed zero-initialized accumulator, optional dbias) are XLA-managed donated buffers — no manual synchronization. Under tracing the offsets values cannot be host-validated, and the weight/scale buffers behind the pointer arrays must stay alive and unmoved across every execution of the traced computation.
Overview
Discrete Grouped GEMM + dGLU backward fusion: A block-scaled grouped GEMM fused with a dSwiGLU/dGeGLU backward epilogue on NVIDIA Blackwell GPUs (SM100+), designed for MoE workloads where each expert weight/scale lives in a separate allocation.
This API uses per-expert pointer tensors instead of packed (N, K, L) tensors:
b_ptrs: deviceint64tensor of per-expert B pointerssfb_ptrs: deviceint64tensor of per-expert SFB pointers
Groups are contiguous in the M dimension and described by padded_offsets (cumulative aligned end offsets).
This kernel performs:
- Block-scaled grouped GEMM backward core using per-expert pointer inputs
- dGLU backward epilogue (
act_func in {"dswiglu", "dgeglu"}) usingc_tensor,beta,prob, anddprob - Optional quantized output with row/column scale factors
Shapes
- Inputs
A: contiguous activation/gradient tensor across all groups, shape(valid_m, K, 1)B_g: expert-gweight tensor referenced byb_ptrs[g], logical shape(N, K)(or(N, K, 1))C: forward activation tensor consumed by backward epilogue, shape(valid_m, N/2, 1)SFA: scale factor tensor for A, shape(32, 4, ceil(valid_m/128), 4, ceil(ceil(K/sf_vec_size)/4), 1)SFB_g: expert-gB scale tensor referenced bysfb_ptrs[g], shape(32, 4, ceil(N/128), 4, ceil(ceil(K/sf_vec_size)/4), 1)padded_offsets: cumulative sum of aligned group M sizes, shape(L,).valid_m = padded_offsets[-1]alpha: per-group scaling factors, shape(L,)beta: per-group scaling factors forC, shape(L,)prob: per-row gating probabilities, shape(valid_m, 1, 1)dprob: probability gradient output buffer, shape(valid_m, 1, 1). Must be zero-initialized.norm_const: normalization constant for FP8 quantization, shape(1,)
- Outputs
D_row: row-quantized dGLU output, shape(valid_m, N, 1)D_col: column-quantized dGLU output, shape(valid_m, N, 1)dprob: updated in-place probability gradient output, shape(valid_m, 1, 1)SFD_row: row scale factors (when SFD outputs are enabled; wrapper auto-enables this for FP8-input configs), shape(32, 4, ceil(valid_m/128), 4, ceil(ceil(N/sf_vec_size)/4), 1)SFD_col: column scale factors (when SFD outputs are enabled; wrapper auto-enables this for FP8-input configs), shape(32, 4, ceil(N/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1)amax: per-group amax (optional; wrapper provides it whend_dtypeis bf16/fp16), shape(L, 2, 1)
Equations
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 (equations shown for act_func="dswiglu"):
For each epilogue tile, the kernel reads paired activation fragments from C and forms gate/input branches:
gate = beta_g * C_gateup = beta_g * C_upsig = sigmoid(gate)swish = gate * sig
Then:
[d_gate, d_up] are interleaved back into full-width D_row/D_col in 32-column blocks.
For act_func="dgeglu", the derivative math switches to dGeGLU.
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
- Shape:
-
Input tensor B pointers:
b_ptrs(wrapper/class execute)- Shape:
(L,)whereL = num_experts - Dtype:
int64, CUDA device tensor - Each pointer must reference one expert B tensor with logical shape
(N, K)(or(N, K, 1)) and dtypeb_dtype - Expert B layout is controlled by
b_major("k"or"n")
- Shape:
-
Input tensor SFB pointers:
sfb_ptrs(wrapper/class execute)- Shape:
(L,) - Dtype:
int64, CUDA device tensor - Each pointer must reference one expert SFB tensor with shape
(32, 4, ceil(N/128), 4, ceil(ceil(K/sf_vec_size)/4), 1)
- Shape:
-
Input tensor C:
c_tensor(wrapper/class)- Shape:
(valid_m, N/2, 1) - Stride:
(N/2, 1, valid_m*(N/2))- must be N-major - Dtype (
c_dtype):{float32, float16, bfloat16}
- Shape:
-
Output tensor D_row:
d_row_tensor(class) or returned in wrapper dict- Shape:
(valid_m, N, 1) - Stride:
(N, 1, valid_m*N)- must be N-major - Dtype (
d_dtype):- FP4 inputs:
{float16, bfloat16, float32} - FP8 inputs:
{float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}
- FP4 inputs:
- Shape:
-
Output tensor D_col:
d_col_tensor(class) or returned in wrapper dict- Shape:
(valid_m, N, 1) - Stride:
(N, 1, valid_m*N)- must match D_row (N-major) - Dtype: Must match D_row
- Shape:
-
Input tensor prob:
prob_tensor(wrapper/class)- Shape:
(valid_m, 1, 1) - Dtype:
float32
- Shape:
-
Output tensor dprob:
dprob_tensor(wrapper/class)- Shape:
(valid_m, 1, 1) - Dtype:
float32 - Must be zero-initialized before kernel execution
- 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:
- SFD_row (optional):
sfd_row_tensor(wrapper) orsample_sfd_row,sfd_row_tensor(class)- Shape:
(32, 4, ceil(valid_m/128), 4, ceil(ceil(N/sf_vec_size)/4), 1) - Dtype: Must match SFA
- Required when: SFD outputs are enabled
- Shape:
- SFD_col (optional):
sfd_col_tensor(wrapper) orsample_sfd_col,sfd_col_tensor(class)- Shape:
(32, 4, ceil(N/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1) - Dtype: Must match SFA
- Required when: SFD outputs are enabled
- Shape:
- SFA (A scale factor):
-
Group offsets
- padded_offsets: Cumulative sum of aligned group M sizes
- Shape:
(L,) - Dtype:
int32 padded_offsets[-1] = valid_m
- Shape:
- padded_offsets: Cumulative sum of aligned group M sizes
-
Scaling tensors
- alpha: Per-group scaling factors
- Shape:
(L,) - Dtype:
float32
- Shape:
- beta: Per-group scaling for
C- Shape:
(L,) - Dtype:
float32
- Shape:
- amax (optional): Per-group max absolute values
- Shape:
(L, 2, 1) - Dtype:
float32 - If provided, updated in-place; wrapper auto-allocates it when
d_dtype in {bfloat16, float16}
- Shape:
- norm_const (optional): Normalization constant for FP8 quantization
- Shape:
(1,) - Dtype:
float32 - Required when:
sfd_row_tensor/sfd_col_tensorare provided
- 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 in {128, 256}TILE_N = 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
- Allowed values:
{16, 32}. Default:16
-
vector_f32: bool- Enable packed f32 operations
- 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 outputs are enabled
- Default:
False
-
act_func: str- Activation derivative. Valid values:
"dswiglu","dgeglu" - Default:
"dswiglu"
- Activation derivative. Valid values:
-
b_major: str- Expert B layout. Valid values:
"k","n" - FP4 inputs require
"k" - Default:
"k"
- Expert B layout. Valid values:
-
epilogue_op: Optional[str]- Optional epilogue transform after backward math
- Valid values:
None,"none","identity","relu","srelu" - Default:
None
-
CUDA stream (
current_streamin class API and wrapper)
Wrapper-specific Parameters: discrete_grouped_gemm_dswiglu_wrapper_sm100
n: int: Logical full N dimension for expert B / D outputsb_dtype: torch.dtype: Dtype of expert B tensors referenced byb_ptrsd_dtype: torch.dtype: Output D tensor dtype. Default:torch.bfloat16cd_major: str: Major dimension for C and D tensors. Must be"n"
Wrapper Return Values
Returns a TupleDict - a dictionary-like object that also supports tuple unpacking and integer indexing.
Dictionary keys (also tuple unpacking order):
d_row_tensor: Row-quantized dGLU outputd_col_tensor: Column-quantized dGLU outputdprob_tensor: Probability gradient outputamax_tensor: Per-group amax (whend_dtype in {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
DiscreteGroupedGemmDswigluSm100 (constructor)
sample_a,num_experts,b_shape,b_dtype,sample_c,sample_d_row,sample_d_col,sample_sfa,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 b_shapemust be logical(N, K)for one expert (pass logicalK, not packedK/2, for FP4)
- Note:
DiscreteGroupedGemmDswigluSm100.execute
a_tensor,b_ptrs,c_tensor,d_row_tensor,d_col_tensor,sfa_tensor,sfb_ptrs,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. Layouts must match constructor sample descriptors.
Support Surface and Constraints
Layouts and Strides
Amust be K-major- Expert
Blayout is selected byb_major:b_major="k": K-majorb_major="n": N-major (FP8 configs only)
C,D_row, andD_colmust be N-major- All tensors must be 16-byte aligned along the contiguous dimension
Data Types
Input/Weight Types (ab_dtype)
Additional Type Constraints
b_dtypemust matchAdtypeSFA,SFD_row, andSFD_colmust share dtypeD_rowandD_colmust have the same dtypeacc_dtypemust befloat32probanddprobmust befloat32c_dtypemust be one of{float32, float16, bfloat16}sf_dtype=float8_e4m3fnwithsf_vec_size=32is not supported- FP8
ab_dtypewithsf_vec_size=16is not supported - FP4
ab_dtyperequiresb_major="k"
Scale Factor Output Requirements
-
When
sfd_row_tensor/sfd_col_tensorare provided:sfd_row_tensor,sfd_col_tensor, andnorm_const_tensorare all required- These must be provided together (all
Noneor all notNone)
-
amax_tensoris optional:- If provided, it is updated in-place with per-group maxima
- Wrapper auto-allocates it when
d_dtype in {bfloat16, float16}
Tiling and Cluster
mma_tiler_mn[0] = 256enables 2-CTA instructions (use_2cta_instrs=True)mma_tiler_mn[0] = 128uses the non-2CTA instruction path- When
use_2cta_instrs=True:cluster_shape_mn[0]must be divisible by 2 m_alignedmust be divisible bymma_tiler_mn[0]m_alignedmust equalFIX_PAD_SIZE=256
Shapes and Divisibility
Dis produced in 32-column blocks (chooseNdivisible by 64 for standard dGLU layouts)Cis half-width relative toD:shape(C)[1] = N/2padded_offsetslengthLis expert count and must be<= 1024valid_m = padded_offsets[-1]determines actual M sizeb_ptrsandsfb_ptrsmust be CUDAint64tensors with shape(L,)
Environment
- Requires CUDA with SM100+ compute capability (Blackwell GPUs)
Usage Examples
For runnable examples and validation, see:
test/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu.pytest/python/fe_api/grouped_gemm/test_discrete_grouped_gemm_dswiglu_utils.py