Mixture of Experts (MoE) Configuration#
This document provides a detailed explanation of the configuration parameters related to Mixture of Experts (MoE) models in MaxText. These settings control the model architecture, routing mechanisms, and performance optimizations. Default values and parameter definitions are located in src/maxtext/configs/base.yml and are primarily used in src/maxtext/layers/moe.py.
1. Architecture#
MoE Strategy#
MaxText supports both Dropless and Dropping strategies. Please refer to the decision tree below to determine the active strategy.
Figure 1: Decision Logic for MaxText MoE Strategies.
Dropless:
Tokamax Ragged Dot (GMM v2): Enabled by setting
sparse_matmul=True, use_tokamax_gmm=True, use_gmm_v2=True.Tokamax Ragged Dot (GMM v1): Enabled by setting
sparse_matmul=True, use_tokamax_gmm=True.Megablox: Enabled by setting
sparse_matmul=True, use_tokamax_gmm=False, megablox=True.JAX Ragged Dot: Enabled by setting
sparse_matmul=True, use_tokamax_gmm=False, megablox=False.Dense Matmul: Enabled by setting
sparse_matmul=False, capacity_factor=-1.
Dropping:
Dense Matmul: Enabled by setting
sparse_matmul=False, capacity_factor > 0(commonly 1.0 to 1.25).
General Configuration#
num_experts: The total number of routed experts available in the MoE layer.
num_experts_per_tok: The number of experts selected for each token, often referred to as top-k strategy.
shared_experts: The number of experts that are always active for every token, in addition to the routed experts.
base_moe_mlp_dim: The intermediate dimension size for the MLP blocks within the experts.
interleave_moe_layer_step: Defines the frequency of MoE layers in transformers. If set to 1, every layer is an MoE layer. If set to X, an MoE layer appears every X layers.
first_num_dense_layers: The number of initial dense layers before the first MoE layer is introduced.
float32_weight_sum: Controls the accumulation precision of the MoE combine reduction after GMM — the weighted sum of the expert outputs \(E_k(x)\) by their routing weights \(g_k\) (not a sum of model parameters): \(y = \sum_{k=1}^{K} g_k E_k(x)\)
True(default): Casts both operands tofloat32before the combine einsum, accumulates infloat32, then casts back to the modeldtype. Recommended for numerical stability.False: Reduces directly in the modeldtype(e.g.bfloat16), halving the operand size. SetFalseonly for HBM-bound recipes ifbfloat16accumulation does not degrade convergence.
Routing Mechanism#
The router computes affinity logits via \(s = \mathrm{score}(x W_g) + b\), where \(x\) is token representations, \(W_g\) is the gate projection kernel, and \(b\) is optional routing bias.
use_random_routing: If enabled, ignores the gate logits and routes tokens to random experts. This is designed to simulate load balancing for debugging and performance testing purposes.
n_routing_groups and topk_routing_group: Experts are divided into n_routing_groups. The router first selects the top k highest-scoring groups (as topk_routing_group), and then selects experts only from those groups.
routed_bias: If enabled, adds a learnable bias term to the gate logits to facilitate load balancing.
routed_bias_update_rate: Defines the update rate to the routed bias term above. Applicable only to the DeepSeek decoder block. For DeepSeek V4, this enables a specialized, auxiliary-loss-free routing bias mechanism. This implementation utilizes a pure nnx.Variable (MoEBiasVar) instead of a standard nnx.Param, which completely isolates the bias update step from the global model optimizer state. The bias is updated directly at the end of the routing step to balance the token distribution mathematically across experts without compromising language modeling convergence. The update is routed_bias_update_rate * sign(mean_load - load), computed from the expert token counts of the full global batch (summed over num_moe_token_chunks), as in Megatron-LM.
float32_gate_logits (default: False): Runs the MoE router computation before GMM (\(s = \mathrm{score}(x W_g) + b\)) in float32 for numerical stability. Operands (\(x\), \(W_g\)) are stored in weight_dtype and cast to float32 at compute time; emitted logits \(s\) remain float32 for downstream Top-K and load-balancing losses. For gemma4, the router norm and scale are also computed in float32.
Quantization Interaction: Incompatible with
quantize_router_proj=True(rejected at config init, as quantizing \(x W_g\) discards thefloat32cast). To run the router infloat32under quantization, setquantize_router_proj=False.
quantize_router_proj (default: True): Applicable when use_qwix_quantization=True and quantization is set; ignored otherwise. Set to False to exclude the router projection matmul (\(x W_g\)) from quantization. To run the projection in float32 under quantization, pair with float32_gate_logits=True.
DeepSeek V4 Auxiliary-Loss-Free & Sequence-Wise Load Balancing#
MaxText implements an exact, paper-aligned version of DeepSeek V4’s load balancing strategies (as specified in the DeepSeek-V4 technical report). The architecture employs two distinct mechanisms:
Auxiliary-Loss-Free Strategy: Handled via
MoEBiasVar, this implementation utilizes a purennx.Variableinstead of a standardnnx.Param. This strictly isolates the bias update step from the global model optimizer state. Unlike the Hugging Face reference implementation—which deviates from the paper by keeping the bias parameters coupled to the global optimizer—our implementation ensures the routing bias balances token distribution mathematically across experts without polluting the main model gradients.Sequence-Wise Balance Loss: An augmenting auxiliary loss (
load_balance_loss_weight) applied to prevent extreme routing imbalance within individual sequences.
routed_score_func: Defines the scoring function for the router.
routed_scaling_factor: A scalar multiplier applied to the expert weights.
load_balance_loss_weight: Sets the coefficient for the auxiliary loss term used to encourage balanced token distribution among experts. By default (moe_use_megatron_seq_aux_loss=False), routers use a Switch Transformer-style loss averaged over MoE layers. When moe_use_megatron_seq_aux_loss=True, sigmoid routers (routed_score_func="sigmoid", e.g. DeepSeek-V3) use Megatron-LM’s sequence-wise aux loss (seq_aux_loss): the probabilities are the float32 sigmoid scores normalized per token, the token fractions come from a top-k without the routing bias or group limit, the loss is computed per full sequence, and the per-layer losses are summed over MoE layers.
moe_use_megatron_seq_aux_loss (default: False): Selects whether sigmoid routers (routed_score_func="sigmoid", with te_moe_block=False) use Megatron-LM’s sequence-wise auxiliary load-balancing loss summed across MoE layers (True) or the Switch Transformer-style auxiliary loss averaged across MoE layers (False).
norm_topk_prob: If enabled, normalizes the router weights for the selected top-k experts.
MLP Block & Computation#
sparse_matmul: Determines whether to use efficient sparse matrix multiplication or dense matrix multiplication.
True: Uses specialized kernels (like Tokamax Ragged Dot or Megablox) or JAX Ragged Dot to perform computation only on active tokens. This is generally faster for MoE.False: Performs dense computation with masking. This is typically used when checking numerical correctness or implementing dropping strategies.
use_tokamax_gmm: If enabled, use Tokamax library’s Ragged Dot for matmul. Recommended for dropless configurations.
use_gmm_v2: If enabled, use the Tokamax GMM v2 kernel for grouped matrix multiplication. Requires use_tokamax_gmm to be True.
use_gmm_v2_heuristic_tiling: If enabled, use the heuristic tiling from Tokamax GMM v2. Recommended when not using custom tuned tile sizes.
megablox: If enabled, use Megablox for sparse matrix operations. Effective only when use_tokamax_gmm is False.
capacity_factor: A scalar multiplier for expert capacity. Effective only when sparse_matmul is False.
Value > 0: Enforces a strict capacity limit; tokens exceeding this limit are dropped.
Value = -1: Dropless with dense matrix multiplication, which is computationally expensive and typically used only as a baseline.
ragged_buffer_factor: A scalar multiplier for the size of the ragged buffer (effectively expert capacity). Effective only when sparse_matmul is True.
Value > 0: Uses an explicit buffer size which may drop tokens when this size is exceeded
Value = -1: Uses a worst case calculated buffer size which is guaranteed to not drop any tokens.
eval_ragged_buffer_factor: Evaluation override for ragged_buffer_factor.
moe_dropless_fallback: What to do when ragged_buffer_factor > 0 and the ragged buffer would drop tokens. This allows training with a smaller tuned buffer for speed/memory efficiency while guaranteeing no token dropping. Requires use_ring_of_experts=True and use_ragged_sort=True.
None(default): tokens that overflow the buffer are dropped.step: the step is run with the capped buffer; if any layer overflows, the step discards its own update in-graph and the entire train step is replayed with a worst-case (dropless) buffer from the state it returned. State donation stays enabled (input/output aliased), since the replay never needs a second copy of the state, and the host checks for overflow after every step. Not supported withenable_dilocoorcompiled_trainstep_file, neither of which has a dropless replay executable, nor withoptimizer_memory_host_offloadorparameter_memory_host_offload.layer: MoE layer calculates from the router’s top-k whether its ragged buffer would overflow, then useslax.condto run either the normal capped path or a chunked dropless path for that layer only. Nothing is replayed, state donation stays enabled, and there is no per-step host sync.
moe_log_max_load_ratio (debug, default False): Logs learning/moe_max_load_ratio and learning/moe_max_load_ratio_mean, the max and mean over MoE layers of the largest expert-shard load divided by the perfectly balanced load. A step overflows the ragged buffer when the max exceeds ragged_buffer_factor, so these metrics help size it. Requires use_ring_of_experts=True and use_ragged_sort=True.
use_custom_sort_vjp: If enabled, use a custom Vector-Jacobian Product (VJP) sort for efficient backward pass processing in sparse matmul. Recommended to replace the inefficient scatter-add generated by the jax.numpy.take in the backward pass.
mlp_bias: If enabled, add learnable bias terms for MLP matmul. Originally implemented to support the GPT-OSS model architecture.
prefuse_moe_weights: If enabled alongside sparse_matmul=True, fuses the two FFN1 grouped GEMMs (wi_0 and wi_1) into a single grouped GEMM call. Expert weights are stored in a concatenated (num_experts, embed_dim, 2 * mlp_dim) shape, so input activations are loaded from HBM once per forward pass instead of twice. Backend-agnostic (works with Megablox, JAX Ragged Dot, and Tokamax). When used with attention=vllm_rpa, the fused weight tensor is passed directly to the vLLM-TPU serving kernel without splitting.
use_batch_split_schedule (experimental): If enabled, split batch into micro-batches to hide communications that yields performance benefits.
TransformerEngine MoEBlock#
te_moe_block: If enabled, uses TransformerEngine’s fused expert-parallel MoEBlock for routing, dispatch, grouped GEMMs, and combining expert outputs. It requires sparse_matmul=True, prefuse_moe_weights=True, and TransformerEngine JAX with expert-parallel MoE support (TE 2.19+).
te_gmm_quantization: Selects the quantization mode for TransformerEngine grouped GEMMs. This must be set when te_moe_block=True. Available options are:
te_no_quant: Uses the model’s default precision, such as BF16, without quantization.te_mxfp8: Uses MXFP8 quantization.
te_ep_overflow_check_every_n_steps: Sets the number of training steps between host-side checks of buffered receive-capacity overflow results when ragged_buffer_factor limits TransformerEngine’s receive capacity. An overflowing step skips its optimizer update immediately on device. At the next check, training raises an error that reports the observed demand and configured capacity. The default is 20.
3. Performance Tuning#
These parameters provide granular control over the tiling dimensions for sparse matmul Pallas kernel.
wi_tile_...: Tile size for the first layer of the MLP (Input -> Hidden).wo_tile_...: Tile size for the second layer of the MLP (Hidden -> Output).
For each, you can control:
..._fwd_...: Tile size for the forward pass...._dlhs_...: Tile size for the backward pass gradient calculation w.r.t. activations...._drhs_...: Tile size for the backward pass gradient calculation w.r.t. weights.
For each dimension, you can control:
..._batch_seq: Tile size for batch x sequence dimension...._embed_dim: Tile size for embedding dimension...._mlp_dim: Tile size for MLP dimension.
Evaluation-Stage Forward Tiling#
Separate evaluation tile sizes are enabled when eval_step uses a custom logical mesh or sharding rule (custom_mesh_and_rule_for_eval=True, where logical_axis_rules_for_eval != logical_axis_rules). MaxText checks the active logical axis rules via max_utils.is_eval(config) to apply the evaluation tile sizes.
When the logical mesh rules are identical between training and evaluation, the training tile sizes are used directly.
Available forward-pass evaluation tile parameters in GMM:
eval_wi_tile_fwd_batch_seq,eval_wi_tile_fwd_embed_dim,eval_wi_tile_fwd_mlp_dimeval_wo_tile_fwd_batch_seq,eval_wo_tile_fwd_embed_dim,eval_wo_tile_fwd_mlp_dim
All eval_* tile sizes default to None. Only forward-pass tile configurations are needed for evaluation since backward gradients (dlhs and drhs) are not computed during evaluation.
Implementation Support:
JAX Ragged Dot:
Supports forward pass only (6 configs:
wi_tile_fwd...andwo_tile_fwd_..., plus theireval_...counterparts).Configs are enabled for INT8, FP8, and BF16.
Megablox:
Supports all 18 configurations (plus the 6
eval_...forward counterparts).Configs are enabled for INT8, FP8, and BF16.
Tokamax Ragged Dot (Includes two implementations):
GMM v1: Uses Tokamax’s native autotuner; does not accept manual tile sizes from MaxText.
GMM v2: Supports all 18 manual tiling configurations (plus the 6
eval_...forward counterparts). Optionally, useuse_gmm_v2_heuristic_tiling=Truefor heuristic tiling.Enabled for FP8 and BF16.