A mixture-of-experts (MoE) layer replaces one large feed-forward block with many smaller ones and a router that sends each token to a few of them. On paper the saving is simple: a token touches k of E experts, so it costs k/E of the FLOPs a dense layer with the same total parameters would. On a GPU the saving is not automatic. A dense feed-forward layer is two or three large matrix multiplies that run near peak. An MoE layer is a sequence of small, data-dependent kernels: a router, a sort, a gather, differently sized matrix multiplies, and a scatter-reduce.

This article walks through the kernels of one MoE layer in the order they run, explains which are bound by memory bandwidth and which by math, shows the index arithmetic that turns ragged routing into tiles a tensor core can consume, and works through the numbers for a Mixtral-shaped layer in decode and prefill. The routing math and load-balancing losses live in MoE routing math, and dropless dispatch strategies in dropless MoE; here the focus is what the GPU executes and how to make it fast.

The kernel inventory of one layer

Take a batch of T tokens with hidden size H, E experts of intermediate size F, and top-k routing. A straightforward implementation launches six kinds of work:

Kernels launched by one MoE layer (forward, single GPU)Hidden statesx: T tokens x H1. RouterGEMM + softmax + top-k2. Aligncount, pad, sort ids3. Permutegather rows4. Expert GEMM 1gate + up, fused SwiGLU5. Expert GEMM 2down, x routing weight6. Combinesum k rows per tokenWith expert parallelism: all-to-all after 3 and before 6dispatch and combine then cross NVLink or the network
The six stages of a single-GPU MoE forward pass. Stages 4 and 5 are grouped matrix multiplies over every expert at once, not E separate launches.
StageWorkUsually bound by
1. Routerx times W_g (H x E), softmax, top-k, renormalizeMemory and launch overhead; E is small
2. AlignCount tokens per expert, prefix-sum, pad, write sorted assignment idsLatency; a few microseconds that matter at small T
3. PermuteGather T*k rows of x into expert-contiguous orderMemory bandwidth
4. Expert GEMM 1Rows times W_gate and W_up (H x 2F), then SwiGLUWeights at small T, math at large T
5. Expert GEMM 2Rows times W_down (F x H), optionally times routing weightSame as stage 4
6. CombineSum the k weighted outputs of each token back into placeMemory bandwidth

The naive PyTorch loop, for e in range(E): y[idx_e] += w_e * ffn_e(x[idx_e]), is the right reference for testing, but at E = 64 or more it launches hundreds of tiny kernels per layer and leaves the GPU idle.

Router: one fused kernel

The router multiplies each token by a gate matrix of shape H x E, applies a softmax (or a sigmoid in some newer models), picks the top k experts and, in many models, renormalizes the k chosen weights so they sum to one. With E between 8 and 256 the GEMM is skinny and cheap; the cost is the separate launches for softmax, top-k and renormalization, each of which reads and writes a T x E tensor. Fusing all of it into one kernel where a warp owns a row, keeps the E logits in registers, and writes only the k ids and k weights removes three round trips to memory. vLLM, for example, ships a fused top-k softmax kernel for this step.

Compute the softmax in FP32 even with BF16 activations: a BF16 softmax produces ties that top-k breaks differently on different kernels. Keep the tie-break rule explicit, such as lowest expert id wins, so the reference and fused kernels agree in tests.

Alignment: making ragged routing tileable

A tensor-core GEMM works on tiles of BLOCK_M rows. The rows for one expert must therefore be contiguous, and a tile must never mix two experts, because each tile multiplies by exactly one expert's weights. The alignment step builds that layout without moving any activation data: it computes, for every assignment (token t, slot j), where that row will sit in an expert-sorted, block-padded order.

Alignment: each expert's rows are padded to a multiple of BLOCK_Mblock -> expert 0block -> expert 0block -> expert 1block -> expert 2block -> expert 2block -> expert 2BLOCK_M = 4: 16 assigned rows become 24 computed rows; dashed cells are padding
Padding to BLOCK_M makes every tile single-expert. The waste is at most BLOCK_M - 1 rows per expert, which is negligible in prefill and dominant in decode.

vLLM calls this function moe_align_block_size; it returns sorted assignment ids, the expert id of every block, and the padded length, and its fused Triton MoE kernel consumes them. A reference version in plain PyTorch shows exactly what it computes:

import torch

def moe_align(topk_ids: torch.Tensor, num_experts: int, block_m: int):
    """topk_ids: [T, k] expert ids. Assignment a = t * k + j.
    Returns padded sorted assignment ids, the expert of every block, and the padded length."""
    flat = topk_ids.flatten()
    counts = torch.bincount(flat, minlength=num_experts)
    padded = (counts + block_m - 1) // block_m * block_m
    starts = torch.cumsum(padded, 0) - padded            # padded start of each expert
    total = int(padded.sum())
    sentinel = flat.numel()                              # ids >= T*k mark padding rows
    sorted_ids = torch.full((total,), sentinel, dtype=torch.int32, device=flat.device)
    order = torch.argsort(flat, stable=True)             # assignments grouped by expert
    exp_sorted = flat[order]
    group_start = (torch.cumsum(counts, 0) - counts)[exp_sorted]
    rank = torch.arange(flat.numel(), device=flat.device) - group_start
    sorted_ids[starts[exp_sorted] + rank] = order.to(torch.int32)
    block_expert = torch.repeat_interleave(
        torch.arange(num_experts, device=flat.device), padded // block_m)
    return sorted_ids, block_expert, total

The production kernel does this in one launch with a shared-memory histogram and a prefix sum. Note int(padded.sum()): it forces a device-to-host sync and breaks CUDA graph capture. Fast paths allocate for the worst case, T*k + E*(BLOCK_M-1) rows, and let the GEMM read the true length from device memory. DeepSeek's DeepGEMM library offers a masked grouped-GEMM layout for exactly this decode case, alongside a contiguous layout for prefill and training.

The fused expert GEMM

The expert GEMM is where the time goes once T is large. The trick that makes it one launch is that every program instance looks up its expert from the block table and its rows from the sorted ids, so the gather of stage 3 can be folded into the GEMM's A-operand load instead of being a separate kernel. In pseudocode for one program that owns block b and output columns n0 to n0 + BLOCK_N:

# one program instance: (block b, column tile n0)            -- Triton-style pseudocode
e    = block_expert[b]                                       # one expert per tile
ids  = sorted_ids[b*BLOCK_M : (b+1)*BLOCK_M]                 # assignment ids, padded
ok   = ids < num_assignments                                 # padding rows are masked
rows = ids // top_k          # token row in x (fused gather); in GEMM 2 the input is
                             # already in assignment order, so there rows = ids
acc  = zeros(BLOCK_M, BLOCK_N, fp32)
for k0 in range(0, H, BLOCK_K):
    a = load(x[rows, k0:k0+BLOCK_K], mask=ok)                # scattered rows of activations
    w = load(W[e, k0:k0+BLOCK_K, n0:n0+BLOCK_N])            # dense tile of one expert
    acc += dot(a, w)
if MUL_ROUTING_WEIGHT:                                       # usually only in the down GEMM
    acc *= topk_weights.flatten()[ids][:, None]
store(out[ids, n0:n0+BLOCK_N], acc.to(bf16), mask=ok)

Two further fusions are common. Gate and up projections are stored as one H x 2F matrix so one GEMM produces both, with the SwiGLU, silu(g) * u, in the epilogue or a small elementwise kernel. The routing weight multiply moves into the down GEMM's epilogue, so combine becomes a plain sum. For the programming model behind kernels like this, see Triton, in depth.

Tile shape is the main tuning knob, and it depends on tokens per expert, not on T. With a few rows per expert, BLOCK_M = 16 wastes little padding; with hundreds, 64 or 128 uses the tensor cores better. vLLM keeps per-shape, per-device tuning tables keyed by batch size for this reason.

Worked example: decode versus prefill

Take a Mixtral-8x7B-shaped layer: H = 4096, F = 14336, E = 8, top-2, BF16 weights. One expert holds 3 x 4096 x 14336 = 176.2 million parameters, or 352 MB. Assume an H100 SXM with about 3.35 TB/s of HBM bandwidth and about 989 TFLOPS of dense BF16 tensor throughput.

Arithmetic intensity. An expert that receives m rows does 2 x m x 176.2M FLOPs and must read its 352 MB of weights once. That is m FLOPs per byte, ignoring activations. The GPU's ratio of math to bandwidth is roughly 989 / 3.35, about 295 FLOPs per byte, so an expert needs around 300 rows before its GEMM stops being bound by weight reads. With top-2 of 8 and perfect balance, that is about 1,200 tokens in the batch.

Decode, batch 1. Two experts are touched: 705 MB of weights per layer, about 0.21 ms at full bandwidth. The tensor cores are almost idle; the layer is a memory copy.

Decode, batch 16. A given expert is missed by one token with probability 6/8, so it is missed by all 16 with probability 0.75 to the 16th, about 1 percent. Essentially all eight experts are read: 2.8 GB, about 0.84 ms per layer, four times the batch-1 cost, while each expert sees only about 4 rows. A dense model's decode step barely changes between batch 1 and 16; an MoE step grows until every expert is touched. This is why MoE serving pushes for large decode batches and expert parallelism, which spreads those weight reads over more GPUs' bandwidth (see MoE all-to-all).

Padding. At batch 16 with BLOCK_M = 64, each expert's 4 rows become a 64-row tile, 16x the compute, yet nearly free because the kernel waits on weights anyway. At batch 4,096 in prefill, padding adds at most 63 rows to about 1,000 per expert, roughly 6 percent of real compute.

Combine, backward and low precision

Combine. A scatter-add with atomics is simple, but addition order varies run to run. A gather-based combine, where each program owns a token and sums its k rows in a fixed order, is deterministic and usually as fast.

Backward. Training needs two grouped GEMMs per projection. The input gradient, dY times W transposed, has the same ragged-M shape as the forward pass. The weight gradient, X transposed times dY, reduces over each expert's variable token count, so the ragged dimension becomes K, which needs a different grouped-GEMM variant and is where small experts lose the most efficiency. Keep the sorted ids from the forward pass to replay the permutation in reverse rather than re-sorting.

Low precision. FP8 expert weights halve the decode weight traffic, which is a direct speedup in the memory-bound regime. Per-tensor scales lose too much accuracy for many models, so recent recipes use fine-grained scales; DeepSeek-V3's report describes 1 x 128 tiles for activations and 128 x 128 blocks for weights. The quantization of activations then belongs in the permute kernel, so rows are converted once while they are being moved anyway.

Profiling and tuning

Start with a timeline, then drill into the expert kernels:

# timeline of a few steps, with NVTX ranges around each MoE stage
nsys profile -t cuda,nvtx -o moe_decode python bench_moe.py --batch 16 --steps 20

# per-kernel detail for the fused expert GEMM only
ncu --set full -k regex:moe -c 20 -o moe_gemm python bench_moe.py --batch 16 --steps 3

Then compare each kernel with its own roofline, not with peak FLOPs. For a memory-bound expert GEMM, achieved bandwidth equals the expert-weight bytes touched divided by kernel time; at batch 16 in the example above, a kernel taking 1.2 ms per layer is reaching about 70 percent of peak bandwidth. At small batch, router, align and combine are dominated by launch overhead, where CUDA graphs and fusion beat kernel tuning. Record tokens-per-expert histograms with timings: the busiest expert sets the grouped GEMM's time, which MoE load balancing addresses.

Failure modes

  • Host syncs in the hot path. Reading the padded length or per-expert counts on the CPU stalls the stream and breaks CUDA graphs. Symptom: gaps between kernels in the nsys timeline that grow with layer count.
  • Wrong sentinel handling. If padding ids are not masked on both load and store, padding rows read out of bounds or overwrite real rows. Test with counts of 0, 1 and exactly BLOCK_M per expert.
  • Index overflow. ids * H in 32-bit arithmetic overflows for large prefill batches with large hidden sizes; compute offsets in 64-bit.
  • Routing nondeterminism. A BF16 softmax or an unstable sort sends tied tokens to different experts on different kernels, which shows up as a fused-versus-reference mismatch far larger than rounding error.
  • Configs tuned for one regime. A tile size picked on prefill benchmarks can slow decode; benchmark the batch sizes you actually serve.

Trade-offs

ChoiceGainsCosts
Fused gather into GEMM loadRemoves a full read and write of T*k rowsScattered A-operand loads; harder to use TMA-style bulk copies
Separate permute kernelDense, contiguous GEMM input; simpler kernelExtra memory traffic and a launch
Padding to BLOCK_MSingle-expert tiles, standard GEMM inner loopUp to BLOCK_M - 1 wasted rows per expert
Block-sparse formulationNo padding waste at the expert levelMore complex kernels and metadata
Atomic combineSimple, overlaps with GEMMNondeterministic summation order
FP8 block-scaled expertsHalf the weight bytes in decodeCalibration, scale handling, accuracy checks

What to do next

  1. Write the naive per-expert loop as a reference and a test that compares it with your fused path for skewed, empty and single-token routings.
  2. Run the reference moe_align above on a few routings by hand and check the padded length against sum(ceil(count / BLOCK_M)) * BLOCK_M.
  3. Compute bytes and FLOPs per expert for your own model shape and find the tokens-per-expert ridge point on your GPU.
  4. Profile decode at batch 1, 16 and 64 with nsys, and record achieved weight bandwidth for the expert GEMMs.
  5. Check the hot path for host syncs and confirm the layer captures in a CUDA graph.
  6. Tune tile configs separately for decode and prefill batch sizes, and log tokens-per-expert histograms in production.
Key takeaway: An MoE layer is six kernels, not one: route, align, permute, two grouped expert GEMMs, and combine. Alignment pads each expert to whole tiles so one launch can serve every expert, and the gather and routing-weight multiply fold into the GEMM. In decode the layer is bound by reading every touched expert's weights, so batch size, expert parallelism and FP8 matter more than tile tuning; in prefill it becomes a real GEMM problem where padding and imbalance cost compute. Measure each kernel against its own roofline and tune both regimes separately.