Grouped GEMM + Quant — Unified (SM100)
Grouped GEMM + Quant — Unified (SM100)
This is an experimental API and subject to change.
Overview
Unified Grouped GEMM + Quant fusion: A block-scaled grouped GEMM with output quantization and per-row gating on NVIDIA Blackwell GPUs (SM100+), designed for MoE (Mixture of Experts) workloads. Implemented with CUTLASS/CUTE. Used for FC2 (forward down-projection) and dFC1 (backward FC1 GEMMs).
This kernel uses the unified BlockScaledMoEGroupedGemmQuantKernel which supports the MoEWeightMode abstraction:
- Dense mode (
MoEWeightMode.DENSE): all expert weights packed into a single contiguous(N, K, L)tensor - Discrete mode (
MoEWeightMode.DISCRETE): each expert weight and scale-factor tensor provided through per-expert device pointer arrays
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
- Per-row gating: Multiplies output by per-row gating probability
- Optional quantized output: Produces row and column scale factors for downstream quantization
Shapes
- 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 tensors, each with shape(N, K), passed viab_ptrsSFA: 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 scale factor tensors, passed viasfb_ptrspadded_offsets: cumulative sum of aligned group M sizes, shape(L,).valid_m = padded_offsets[-1]alpha: per-group scaling factors, shape(L,)prob: per-row gating probabilities, shape(valid_m, 1, 1). Required.norm_const: normalization constant for FP8 quantization, shape(1,)
- Outputs
D: row-quantized output, shape(valid_m, N, 1)D_col: column-quantized output, shape(valid_m, N, 1)SFD_row: row scale factors (when SFD outputs are enabled), 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), shape(32, 4, ceil(N/128), 4, ceil(ceil(valid_m/sf_vec_size)/4), 1)amax: per-group amax (whend_dtypeis bf16/float16), shape(L, 1)
Equations
Step 1: Block-scaled grouped GEMM (per group g with rows m in [padded_offsets[g-1], padded_offsets[g])):
Step 2: Per-row gating:
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/sample_a- Shape:
(valid_m, K, 1), Stride: K-major - Dtype:
{float4_e2m1fn_x2, uint8, float8_e4m3fn, float8_e5m2}
- Shape:
-
Input tensor B:
b_tensor/sample_b- Shape:
(N, K, L), Stride: K-major (FP8 also supports N-major) - Dtype: Must match A
- Shape:
-
Output tensor D:
d_tensor/sample_d- Shape:
(valid_m, N, 1), Stride: N-major - Dtype:
{float16, bfloat16, float32}for FP4;{float16, bfloat16, float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}for FP8
- Shape:
-
Output tensor D_col:
d_col_tensor/sample_d_col- Shape/Dtype: Must match D
- Wrapper behavior: returned only when
d_dtype ∈ {float8_e4m3fn, float8_e5m2, float4_e2m1fn_x2}; forbfloat16,float16, andfloat32,outputs["d_col_tensor"]isNone
-
Input tensor prob:
prob_tensor/sample_prob- Shape:
(valid_m, 1, 1), dtype:float32 - Required: pass ones tensor if no gating needed
- Shape:
-
Scale factor tensors: SFA, SFB, SFD_row, SFD_col — block-scaled 6-D layout
-
Group offsets:
padded_offsetsshape(L,), dtypeint32 -
Scaling tensors:
alphashape(L,),amaxshape(L, 1),norm_constshape(1,)
Common Parameters
acc_dtype: Must betorch.float32mma_tiler_mn: Default(256, 256); supported tiles areTILE_M ∈ {128, 256}andTILE_N = 256cluster_shape_mn: Default(2, 1)whenTILE_M=256,(1, 1)otherwisesf_vec_size:{16, 32}. Default:16vector_f32: Default:Falsem_aligned: Must be256discrete_col_sfd: Default:False
Wrapper Return Values
Returns TupleDict: d_tensor, d_col_tensor (optional; None for bfloat16/float16/float32 outputs), amax_tensor, sfd_row_tensor, sfd_col_tensor
Support Surface and Constraints
Data Types
Key Constraints
AandBmust have same dtype;DandD_colmust have same dtype- All scale factor tensors must have same dtype
- Expert count
<= 1024; M aligned to 256 - SM100+ compute capability required
prob_tensoris unconditionally requireduse_single_group_runtime_offsets=Truerequires exactly one expert. 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,).
Usage Examples
For usage examples, see test/python/fe_api/grouped_gemm/test_grouped_gemm_quant.py + test/python/fe_api/grouped_gemm/test_grouped_gemm_quant_utils.py (dense and discrete unified API coverage)