grouped-gemm-moe

A variable-size batched expert GEMM with the Mixture-of-Experts token permutation fused into the matmul, trainable end to end, loadable through kernels. The reference baselines are torch._grouped_mm, a pure-torch MoE (gather, per-expert matmul, scatter, combine), and PyTorch autograd, matched to a few parts in ten thousand with the fp32 fusion exact bit-for-bit.

A MoE layer routes each token to a few of E experts, so the work is a batch of GEMMs with variable row counts wrapped by a gather and a scatter. Done directly, both wrappers materialize [num_assignments, hidden] tensors, pure memory traffic, and padding to equal group sizes wastes compute. This kernel gathers rows as they are read into shared memory and scatters results with the routing-weight combine as they are written, so neither intermediate exists, the variable group sizes are respected without padding, an expert with no tokens costs nothing, and the backward carries the same fusion, so gradients flow to tokens, expert weights, and routing weights.

Token particles route into eight variable-width expert bins, one left empty and free, then the grouped GEMM sweeps the tiles

4,096 tokens routing top-2 into eight experts, real assignment counts per tile and one expert deliberately empty at zero cost; the fused path matches the eager MoE to 1e-6 and runs the bf16 forward in 0.72 ms against 1.24 ms eager, with no [assignments, hidden] intermediate ever materialized.

Usage

import torch
from kernels import get_kernel

moe = get_kernel("phanerozoic/grouped-gemm-moe", version=1, trust_remote_code=True)

T, K, N, E, k = 4096, 2048, 2048, 8, 2
x = torch.randn(T, K, device="cuda", requires_grad=True)        # tokens
W = torch.randn(E, N, K, device="cuda", requires_grad=True)     # expert weights
expert_ids = torch.stack([torch.randperm(E, device="cuda")[:k] for _ in range(T)])
topk_weights = torch.rand(T, k, device="cuda", requires_grad=True)
y = moe.moe(x, W, expert_ids, topk_weights)                     # [T, N]
y.sum().backward()                          # grads to x, W, topk_weights

# bare variable-size grouped GEMM (rows already grouped by expert)
out = moe.grouped_gemm(a, W[:4].detach(), group_sizes)

version selects the release branch; trust_remote_code is required by kernels for publishers without the trusted-publisher mark.

API

Symbol Purpose
moe(x, weight, expert_ids, topk_weights) fused MoE, differentiable wrt x, weight, topk_weights
grouped_gemm(a, b, group_sizes) bare variable-size grouped GEMM, matches torch._grouped_mm
moe_forward(x, weight, row2token, topk_weight, group_sizes, T) low-level fused forward over a pre-sorted routing
moe_dweight(A, B, group_sizes) per-expert weight-gradient reduction dW[e] = A[e]^T B[e]
route_and_sort(expert_ids, topk_weights, E) top-k routing -> (row2token, topk_weight, group_sizes)

Method

The routing becomes a sort by expert; from the group sizes the launcher lays out one row-tile per BM rows of each expert. Per tile the GEMM kernel gathers activation rows through the row -> token map straight into shared memory, multiplies against that tile's expert weight with fp32 accumulation (a 128x128 mma.sync grouped GEMM with a cp.async pipeline for bf16, a register-blocked tile loop for the exact fp32 path), and scatters each result into its original token's output row scaled by the routing weight with an atomic add. The backward reuses the same machinery: the token gradient is the fused forward on the incoming gradient with transposed expert weights; the weight gradient is a per-expert A^T B reduction; the routing-weight gradient is an inner product with the unweighted product.

Measured

  • The variable-size grouped GEMM agrees with torch._grouped_mm to 2e-5 relative Frobenius (bf16) and an independent per-group reference to 4e-3 max relative, including uneven and empty groups and non-tile-multiple dimensions.
  • The fused dispatch matches a pure-torch MoE to 2.8e-4 (top-1) and 3.8e-4 (top-2) in fp32, 2.3e-3 relative Frobenius in bf16.
  • The fusion is exact: the fused forward equals the explicit gather-grouped-GEMM-scatter path bit-for-bit in fp32 (max abs diff 0).
  • The backward matches autograd to 3e-4 (tokens), 2e-4 (expert weights), 8e-5 (routing weights), and agrees with finite differences.
  • The bf16 tensor-core forward is about 19x the portable tiled path on an H200; measured against an eager bf16 MoE loop at T = 4096, E = 8: 0.72 ms fused vs 1.24 ms eager.

Requirements and limits

  • NVIDIA GPU with compute capability 8.0+ (one source spans 8.0 to 12.0).
  • float32 or bfloat16 activations and weights; contraction and output dimensions arbitrary (tile-boundary guarded).
  • Gating (softmax and top-k selection) is the caller's; the routing sort runs in torch.

References

Gale et al., "MegaBlocks" (block-sparse grouped GEMM without padding); Tan et al., ScatterMoE (fused gather-scatter GEMM); Lepikhin et al., "GShard"; Fedus et al., "Switch Transformer"; torch._grouped_mm.

License

Apache-2.0.

Downloads last month
-
apache-2.0
Supported hardwares new
CUDA
8.08.68.99.010.012.0
GPU
B300
288GB
NVIDIA SXM
B200
192GB
NVIDIA SXM
H200
141GB
NVIDIA SXM
H100
80GB
GPU
H800
80GB
GPU
H20
96GB
GPU
L40s
48GB
GPU
L40
48GB
GPU
L20
48GB
GPU
L4
24GB
DGX Spark
GB10
128GB
GPU
RTX PRO 6000 WS
96GB
GPU
RTX PRO 6000 Max-Q
96GB
GPU
RTX PRO 5000
48GB
GPU
RTX PRO 4500 WS
32GB
GPU
RTX PRO 4000
24GB
GPU
RTX PRO 4000 SFF
24GB
GPU
RTX PRO 2000
16GB
GPU
RTX 6000 Ada
48GB
GPU
RTX 5880 Ada
48GB
RTX
RTX 5000 Ada
32GB
GPU
RTX 4500 Ada
24GB
RTX
RTX 4000 Ada
20GB
RTX
RTX 4000 SFF Ada
20GB
GPU
RTX 3500 Ada Mobile
12GB
GPU
RTX 2000 Ada
16GB
GPU
RTX A6000
48GB
GPU
RTX A5000
8GB
GPU
RTX A5000 Max-Q
16GB
GPU
RTX A5000 Mobile
16GB
GPU
RTX A4000
16GB
GPU
RTX A4000 Max-Q
8GB
GPU
RTX A4000 Mobile
8GB
GPU
RTX A3000 Mobile
6GB
GPU
RTX A2000
6GB
GPU
RTX A2000 Embedded
4GB
GPU
RTX A2000 Max-Q
4GB
GPU
RTX A2000 Mobile
4GB
GPU
A800
40GB
GPU
A100
80GB
GPU
A40
48GB
GPU
A30
24GB
GPU
A10
24GB
GPU
A2
16GB
RTX
RTX 5090
32GB
RTX
RTX 5090 D
32GB
RTX
RTX 5090 Mobile
24GB
RTX
RTX 5080
16GB
RTX
RTX 5080 Mobile
16GB
RTX
RTX 5070
12GB
RTX
RTX 5070 Mobile
8GB
RTX
RTX 5070 Ti
16GB
RTX
RTX 5070 Ti Mobile
12GB
RTX
RTX 5060 Ti
16GB
RTX
RTX 5060
8GB
RTX
RTX 5060 Mobile
8GB
RTX
RTX 5050
8GB
RTX
RTX 5050 Mobile
8GB
RTX
RTX 4090
24GB
RTX
RTX 4090D
24GB
RTX
RTX 4090 Mobile
16GB
RTX
RTX 4080 SUPER
16GB
RTX
RTX 4080
16GB
RTX
RTX 4080 Mobile
12GB
RTX
RTX 4070
12GB
RTX
RTX 4070 Mobile
8GB
RTX
RTX 4070 Ti
12GB
RTX
RTX 4070 Super
12GB
RTX
RTX 4070 Ti Super
16GB
RTX
RTX 4060
8GB
RTX
RTX 4060 Ti
8GB
RTX
RTX 4090 Laptop
16GB
RTX
RTX 4080 Laptop
12GB
RTX
RTX 4070 Laptop
8GB
RTX
RTX 4060 Laptop
8GB
RTX
RTX 4050 Laptop
6GB
RTX
RTX 3090
24GB
RTX
RTX 3090 Ti
24GB
RTX
RTX 3080
12GB
RTX
RTX 3080 Ti
12GB
RTX
RTX 3080 Mobile
16GB
RTX
RTX 3070
8GB
RTX
RTX 3070 Ti
8GB
RTX
RTX 3070 Ti Mobile
8GB
RTX
RTX 3060 Ti
8GB
RTX
RTX 3060
12GB
RTX
RTX 3060 Mobile
6GB
RTX
RTX 3050 Mobile
4GB
GPU
RTX 2050 Mobile
4GB
Jetson
Jetson AGX Orin 64GB
64GB
Jetson
Jetson AGX Orin 32GB
32GB
Jetson
Jetson Orin NX 16GB
16GB
Jetson
Jetson Orin NX 8GB
8GB
Jetson
Jetson Orin Nano 8GB
8GB
Jetson
Jetson Orin Nano 4GB
4GB
OS
linux
Arch
x86_64
Kernel Builder
570dcf4