Fused MoE Kernel Optimization
Overview
Fused MoE combines token routing, expert GEMM, and output accumulation into a single kernel, avoiding expensive scatter/gather operations in global memory.
MoE kernels are compute-bound for large expert dimensions and latency-bound for small token counts per expert. The key metric is TFLOPS.
Core Technique
Token-Expert Routing
- Top-k gating: select k experts per token (typically k=2)
- Sort tokens by assigned expert
- Pad each expert's token count to be divisible by BLOCK_SIZE_M
Grouped GEMM Structure
The kernel is a standard tiled GEMM with an extra indirection:
sorted_token_idsmaps program blocks to actual token rowsexpert_idsmaps each M-block to an expert (selects which weight matrix)- Block-level early exit: skip if
pid_m * BLOCK_M >= num_tokens_post_padded
Optimization Patterns
- L2 cache grouping: same as standard GEMM (GROUP_SIZE_M)
- Persistent cache-aware: for many experts, persistent kernel with tile scheduling
- Compute + comm fusion: overlap expert compute with all-to-all communication
Key Parameters
| Parameter | Typical Values | Notes |
|---|---|---|
| BLOCK_SIZE_M | 16, 32, 64, 128 | 16 for small token counts |
| BLOCK_SIZE_N | 32, 64, 128 | Expert hidden dim |
| BLOCK_SIZE_K | 32, 64, 128 | Input feature dim |
| GROUP_SIZE_M | 1, 4, 8 | L2 grouping |
| top_k | 1, 2 | Experts per token |
Small-M heuristic: when M <= num_experts, use BLOCK_M=16, BLOCK_N=32, BLOCK_K=64, GROUP_M=1.
Verification
python skills/kernels/fused-moe/test_fused_moe.py
- Correctness: compare vs sequential per-expert
torch.matmulwith top-k selection - Performance: TFLOPS =
2 * M * top_k * N * K / latency / 1e12
Common Pitfalls
- Incorrect token ID mapping:
sorted_token_ids // top_kmaps back to original token - Padding tokens: padded IDs must be >= num_valid_tokens for masking
- Expert weight layout: typically
[E, N, K](expert-major), not[E, K, N] - Routed weight multiplication: must apply
topk_weightsto correct output rows