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.

Illustration of MoE 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 to float32 before the combine einsum, accumulates in float32, then casts back to the model dtype. Recommended for numerical stability.

  • False: Reduces directly in the model dtype (e.g. bfloat16), halving the operand size. Set False only for HBM-bound recipes if bfloat16 accumulation 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 the float32 cast). To run the router in float32 under quantization, set quantize_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:

  1. Auxiliary-Loss-Free Strategy: Handled via MoEBiasVar, this implementation utilizes a pure nnx.Variable instead of a standard nnx.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.

  2. 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 with enable_diloco or compiled_trainstep_file, neither of which has a dropless replay executable, nor with optimizer_memory_host_offload or parameter_memory_host_offload.

  • layer: MoE layer calculates from the router’s top-k whether its ragged buffer would overflow, then uses lax.cond to 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.

2. Sharding#

use_ring_of_experts (experimental): This feature requires expert parallelism. If enabled, it replaces the standard two All-to-All communications with All-Gather in dispatch and Reduce-Scatter in collect. By gathering inputs across all shards, it allows for local routing and Top-K calculations, followed by result aggregation via Reduce-Scatter. This approach is particularly effective for models with a large Top-K, as it gathers activations before they are replicated k times to reduce communication.

moe_quantize_token_all_gather: If enabled, quantizes token activations to the quantized dtype (e.g. FP8) prior to the Ring of Experts All-Gather across the EP mesh axis, reducing inter-chip communication volume proportionally to the quantized width (2x for FP8 vs BF16). Requires use_ring_of_experts=True, use_gmm_v2=True, and static activation calibration. See Quantization guide for full pipeline details. Default is False.

moe_topk_before_ep_all_gather: If enabled with use_ring_of_experts=True, each expert shard runs top-k on its local tokens and all-gathers only the resulting (batch, seq, num_experts_per_tok) weights and expert ids, instead of all-gathering the (batch, seq, num_experts) router logits and repeating top-k on the full gathered batch on every shard. This removes the redundant top-k and cuts the router all-gather (and its backward reduce-scatter) by a factor of num_experts / num_experts_per_tok. Only the top-k (the index map) moves before the all-gather; the token sort/permutation still happens after it. Expert assignment is unchanged, and loss and gradients match up to floating-point summation order. Ignored with forced, hash or random routing. Default is False.

moe_fsdp_use_two_stage_all_gather: If enabled, split the All-Gather operation for MoE weights into two separate stages when using FSDP/FSDP-transpose sharding. This is preferred when 3D All-Gather support is unavailable.

MoE FSDP Sharding Strategies (Note: At most one of the following three flags can be enabled at a time):

  • shard_exp_on_fsdp: If enabled, shard the expert dimension of the MLP weights on the FSDP axis. Works for both unquantized and quantized. When quantization and weight_quantization_calibration_method are fixed, it performs quantized weight all gather over fsdp before gmm. This is recommended only when num_experts is a multiple of fsdp_parallelism.

  • use_2d_fsdp_sharding: If enabled, use fsdp and fsdp_transpose axes for sharding the MoE weights.

  • shard_embed_moe_on_fsdp: If enabled, keep embed_moe sharded so we can manually QAG (Quantize-All-Gather) it over FSDP. Requires quantization to be specified and weight_quantization_calibration_method to be fixed.

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_dim

  • eval_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... and wo_tile_fwd_..., plus their eval_... 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, use use_gmm_v2_heuristic_tiling=True for heuristic tiling.

    • Enabled for FP8 and BF16.