A mixture-of-experts layer turns one large matrix multiply into many small ones. After the router picks experts, each expert e receives some number of tokens Me that changes every step, and it multiplies them by its own weight matrix of fixed shape K by N. Grouped GEMM is the kernel family that runs all of those products in a single launch, each with its own M, reading tokens from one contiguous buffer and weights from an expert-indexed array. It is the reason a 64-expert layer can run close to the speed of a dense layer with the same total FLOPs instead of spending most of its time launching and draining tiny kernels.
This article is about the kernel itself: how groups are described, how a persistent tile scheduler walks them, why alignment and tile height matter when experts are unevenly loaded, what the backward pass needs, and how to choose between the implementations you can actually call today. Routing math and token dispatch are covered in MoE routing math and dropless MoE; here we assume tokens already arrive sorted by expert.
Why not a loop or a batched GEMM
There are three obvious ways to run E expert matmuls, and each fails for a different reason. A Python loop issues E separate GEMM launches. With 64 experts and a few hundred tokens each, every launch has only a few dozen output tiles, so it fills a fraction of a GPU with more than a hundred SMs, then waits for its slowest tile before the next launch starts. Launch overhead and the half-empty tail of each kernel dominate.
Batched GEMM (one launch, many problems) requires every problem to have the same shape. To use it you must pad every expert to the largest M in the step, which is exactly the capacity-factor approach: either you drop tokens above a cap or you compute on a mountain of zeros. A block-diagonal or block-sparse formulation treats the whole layer as one sparse matrix product; it works well but needs its own sparse metadata and kernels.
Grouped GEMM keeps the per-problem shapes. The kernel receives an array of problem sizes (or offsets into one buffer) plus per-group pointers, enumerates every output tile across all groups, and distributes those tiles over the SMs as if they were one large GEMM.
| Approach | Launches | Handles variable M | Main cost |
|---|---|---|---|
| Loop of GEMMs | E | Yes | Launch overhead, idle SMs in each small kernel |
| Batched GEMM | 1 | No, pad to max M | Wasted FLOPs or dropped tokens |
| Block-sparse GEMM | 1 | Yes, in blocks | Sparse metadata, block-size padding |
| Grouped GEMM | 1 | Yes | Partial tiles at group edges, scheduling logic |
Problem descriptors and the tile scheduler
Each group g is a problem Cg = Ag Bg with shape Mg x K times K x N. In an MoE forward pass K and N are the same for every expert (hidden size and FFN width), and only M varies. The kernel picks an output tile of BM by BN, so group g contributes ceil(Mg/BM) x ceil(N/BN) tiles. A prefix sum over those counts gives each group a starting tile index, and the total is the length of one flat tile space.
A persistent kernel launches roughly one cooperative thread array (CTA) per SM and lets each CTA loop over tile indices with a stride equal to the number of CTAs. For each tile index it finds the owning group (a binary search over the prefix sums, or a cheap forward walk because indices only increase), converts the local index to a row block and a column block, sets up pointers into A, Bg and C, and runs the usual K loop on tensor cores. Rows past the end of the group are masked on load and on store. CUTLASS's grouped kernels follow this pattern with a problem-visitor or tile-scheduler object, and the Triton grouped GEMM tutorial implements the same loop in a few dozen lines.
# Persistent grouped GEMM, one program per SM (Triton-style pseudocode)
tiles_n = cdiv(N, BN)
tiles[g] = cdiv(M[g], BM) * tiles_n # computed on device from offs
start = exclusive_cumsum(tiles); total = start[E-1] + tiles[E-1]
for tile in range(pid, total, num_programs):
g = upper_bound(start, tile) - 1 # owning expert
local = tile - start[g]
mb, nb = local // tiles_n, local % tiles_n
row0 = (0 if g == 0 else offs[g-1]) + mb * BM
rows = min(BM, M[g] - mb * BM) # ragged last block
acc = zeros(BM, BN, fp32)
for k0 in range(0, K, BK):
a = load(A, row0, k0, mask_rows=rows) # tokens of expert g
b = load(B[g], k0, nb * BN) # weights of expert g
acc = mma(a, b, acc)
store(C, row0, nb * BN, acc.to(bf16), mask_rows=rows)Two details matter for speed. First, ordering: walking tiles column-block-major inside a group lets consecutive CTAs reuse the same Bg tiles from L2. Second, the group sizes live on the device. If the host must read Mg back to compute the tile count, every MoE layer forces a device-to-host sync, which stalls the CPU and breaks CUDA graph capture. Good implementations compute the schedule on the GPU from the offsets tensor.
Data layout, offsets and alignment
Grouped GEMM for MoE expects the dispatched tokens as one 2D tensor of shape (total assignments, K), sorted so that every expert's rows are contiguous, plus a 1D int32 tensor of end offsets. If expert counts are [3, 0, 5], the offsets are [3, 3, 8]: group 0 is rows 0 to 2, group 1 is empty, group 2 is rows 3 to 7. Empty experts must be legal, because they happen in practice.
Alignment is where implementations differ. Hardware copy engines such as Hopper's TMA want 16-byte aligned addresses, and some kernels only accept group sizes that are a multiple of their tile height or of a small constant. Frameworks satisfy this by padding each expert's row count up to a multiple (torchtitan, for example, pads to a small multiple of rows) and writing zeros into the gap. Padding to 16 is cheap; padding to a full 128-row tile is not, as the worked example below shows. Read the alignment rules of the kernel you call rather than guessing, and assert them before the call.
import torch
import torch.nn.functional as F
def expert_ffn_up(x, w_up, expert_idx, num_experts):
"""x: (T*k, K) bf16 token copies; expert_idx: (T*k,) expert per copy;
w_up: (E, N, K) bf16, stored like nn.Linear weights. Returns original order."""
assert x.shape[1] % 8 == 0 and w_up.shape[1] % 8 == 0 # 16-byte rows
order = torch.argsort(expert_idx, stable=True) # group rows by expert
n = x.shape[0]
xs = torch.cat([x[order], x.new_zeros(1, x.shape[1])]) # offs[-1] < rows
counts = torch.bincount(expert_idx, minlength=num_experts)
offs = torch.cumsum(counts, dim=0, dtype=torch.int32) # end offset per expert
mm = getattr(F, "grouped_mm", None) or torch._grouped_mm # public from 2.10
ys = mm(xs, w_up.transpose(-2, -1), offs=offs)[:n]
out = torch.empty_like(ys)
out[order] = ys # unpermute
return out
def reference(x, w_up, expert_idx):
"""Loop version: slow, but the oracle for tests."""
out = x.new_empty(x.shape[0], w_up.shape[1])
for e in range(w_up.shape[0]):
rows = (expert_idx == e).nonzero(as_tuple=True)[0]
out[rows] = (x[rows].float() @ w_up[e].float().T).to(x.dtype)
return outThe PyTorch documentation for torch.nn.functional.grouped_mm describes offs as monotonically increasing int32 offsets where offs[i] marks the end of group i and must be strictly less than the row count, hence the extra zero row. It lists BF16 inputs on GPUs of compute capability 8.0 or higher, and for the forward pass passes weights stored as (E, N, K) transposed. It is documented from release 2.10; earlier releases have the private torch._grouped_mm. Recheck these rules for your version.
The backward pass
Training needs two more products per layer, and they are not the same kind of grouped GEMM. Write the forward pass as Yg = Xg Wg with Xg of shape Mg x K.
- Input gradient. dXg = dYg WgT. This is the forward pattern again: jagged M, fixed K and N, same offsets. The kernel only needs to read the weights transposed.
- Weight gradient. dWg = XgT dYg. Now the output has a fixed shape K x N for every expert, and the jagged dimension Mg is the reduction dimension. Each group contributes the same number of output tiles but a different number of K-loop iterations. PyTorch's grouped_mm covers this case when both operands are 2D and the offsets delimit the shared dimension; CUTLASS and other libraries describe it as a grouped GEMM with variable K.
The weight-gradient case has a trap: an expert that received zero tokens still needs dWg = 0 written to memory. A kernel that simply skips empty groups leaves whatever was in the buffer, and the optimizer then applies garbage to that expert. Test the empty-group case explicitly. Load imbalance also shows up differently: in the forward pass a hot expert has many tiles, which the scheduler spreads across SMs, but in dW a hot expert has the same number of tiles with long K loops. Implementations either split K across CTAs and reduce, or accept that the hot expert's tiles finish last.
Worked example: tile counts under skew
Take one MoE up-projection on a GPU with 132 SMs, a persistent kernel with one CTA per SM, BM = BN = 128, K = 4096 and N = 2048. A micro-batch of 8,192 tokens with top-2 routing produces 16,384 assignments over 64 experts, an average of 256 per expert.
Balanced step. Each expert has 256 rows, so 2 row blocks times 16 column blocks = 32 tiles, 2,048 tiles in total. That is 15.5 waves of 132, rounded up to 16, and every computed row is a real token. Useful work is 2 x 16,384 x 4,096 x 2,048, about 275 GFLOP, about 0.55 ms at an effective 500 TFLOP/s.
Skewed step. Now 32 experts get 64 tokens, 28 get 400 and 4 get 784 (still 16,384 in total). Row blocks per expert are 1, 4 and 7, so the layer has 32 + 112 + 28 = 172 row blocks and 2,752 tiles, which is 21 waves. Computed rows are 172 x 128 = 22,016 against 16,384 real ones: 74 percent of the tensor-core work is useful, and the step costs roughly 1.34 times the balanced one even though the FLOPs requested are identical.
Smaller tiles. With BM = 64 the same skew needs 1, 7 and 13 row blocks: 280 blocks, 17,920 computed rows, 91 percent useful. But a 64-row tile reuses each loaded weight tile half as many times, so it moves more bytes per FLOP and can be slower on a balanced step. This is why libraries ship several tile configurations and why some pick one per call from the observed count distribution.
The loop alternative. A per-expert loop would launch 64 kernels, and each small expert would occupy 16 of 132 SMs for one short wave.
Choosing an implementation
| Option | What it is | When to use it |
|---|---|---|
| cuBLAS grouped batched | cublasGemmGroupedBatchedEx, added in cuBLAS 12.5; variable shapes per group | C++ stacks that already use cuBLAS; NVIDIA's launch blog reported about 1.2x over a looped batched API on its MoE example |
| CUTLASS grouped kernels | Templates with a per-group problem array and tile scheduler, including Hopper variants | Custom epilogues, FP8 scaling, fusing activation into the GEMM |
| PyTorch grouped_mm | Offsets-based op used by torchtitan's MoE | PyTorch training where you want autograd and compile support |
| Triton | Write the persistent loop yourself; the official tutorial is a starting point | Research, odd layouts, fused gating |
| MegaBlocks and DeepGEMM | Block-sparse kernels and FP8 grouped kernels from the MoE community | When their layout and precision match yours |
Whatever you pick, fuse what you can around it. The gated FFN needs two up-projections (gate and up) that can be one grouped GEMM with N doubled, followed by an activation that can live in the epilogue. FP8 training adds per-group or per-block scale factors that must travel with each group; CUTLASS epilogues and tensor core formats determine what is cheap.
Failure modes
- Host sync per layer. Calling
.item()or.tolist()on counts to build the schedule. Symptom: CPU-bound step, CUDA graphs fail to capture. - Wrong offset convention. Passing start offsets where end offsets are expected shifts every group by one expert. Outputs look plausible and loss still falls slowly. Only a comparison with the loop reference catches it.
- Unaligned groups. Group sizes or base addresses that violate the kernel's alignment rule cause errors or silent fallbacks to slow paths. Assert alignment.
- Stale dW for empty experts. Covered above; zero-initialize or have the kernel write zeros.
- Tile config tuned on balanced data. Benchmarks with uniform counts hide the padding waste that real routing produces. Replay recorded count histograms.
- Permutation cost ignored. Sorting, gathering and unpermuting rows can rival the GEMM at small hidden sizes. Profile the whole layer.
What to do next
- Log per-expert token counts for a few hundred training steps and plot their distribution.
- Write the loop reference above and a test that compares it with your grouped kernel, including an empty expert and a single-row expert.
- Measure the useful-row fraction (real rows divided by computed rows) for your tile height on the recorded counts.
- Try at least two tile heights and keep the faster one on recorded, not uniform, counts.
- Check that building the schedule causes no device-to-host sync, then capture the layer in a CUDA graph.
- Verify dW for experts that received zero tokens is exactly zero.
- Read expert parallelism next, because across GPUs the same imbalance reappears as communication skew.