How MoE Expert Load Balancing Works with Quantile Balancing in Marin
Marin uses an auxiliary load‑balancing loss combined with quantile‑based (QB) routing to evenly distribute tokens across Mixture‑of‑Experts (MoE) layers, with the QB router dynamically setting selection thresholds via TOPK or HIST estimators.
MoE expert load balancing with quantile balancing is a critical technique for training stable, high‑performance sparse transformer models. In the marin-community/marin repository, this system combines a per‑layer auxiliary loss that penalizes imbalanced routing with a sophisticated quantile‑based router that adaptively determines expert selection thresholds. Together, these mechanisms ensure efficient token distribution without manual tuning.
The Auxiliary Load‑Balancing Loss
The foundation of MoE load balancing in Marin is an auxiliary loss applied at every MoE layer. This loss encourages the router to assign tokens uniformly across all experts, preventing the "expert collapse" where a few experts dominate.
Loss Formulation
According to the source code in experiments/grug/moe_hero_ep/model.py (lines 439–447), the load‑balancing loss follows the HuggingFace reference implementation:
load_balancing_loss = num_experts * jnp.sum(token_fraction * p)
Where:
token_fraction— The actual fraction of tokens each expert received (capacity used divided by total tokens)p— The mean softmax probability assigned to each expert by the router, averaged across all tokensnum_experts— The total number of experts in the layer
This formulation penalizes correlation between high routing probabilities and high token counts. When an expert receives both a large probability mass and a large token fraction, the loss increases, pushing the router toward more uniform distributions.
Implementation in Training
The loss is computed per layer, summed across all MoE layers, and added to the main training objective. Here's how to manually compute it for debugging or analysis:
import jax.numpy as jnp
def load_balancing_loss(token_fraction, router_probs, num_experts):
# token_fraction: shape [Experts]
# router_probs: shape [Token, Experts]
p = jnp.mean(router_probs, axis=0) # average probability per expert
return num_experts * jnp.sum(token_fraction * p)
# Example with dummy values
token_fraction = jnp.array([0.05, 0.04, 0.06, 0.05]) # per‑expert usage
router_probs = jnp.array([[0.2, 0.3, 0.4, 0.1],
[0.25, 0.35, 0.3, 0.1]])
loss = load_balancing_loss(token_fraction, router_probs, num_experts=4)
Quantile Balancing: Dynamic Threshold Selection
While the load‑balancing loss shapes router behavior over training steps, quantile balancing (QB) determines the immediate expert selection threshold during the forward pass. The QB router in Marin computes a dynamic threshold β based on the distribution of logit margins.
The QB Routing Formula
The router computes margins as score − α, where score is the raw expert logit and α is a learned baseline. Experts are selected when their margin exceeds β, where β is the (1 − K/E) upper quantile of the margin distribution (K = experts per token, E = total experts).
Two Quantile Estimators
Marin provides two strategies for estimating this quantile, configured via GrugModelConfig.qb_estimator (lines 40–48 in experiments/grug/moe_hero_ep/model.py):
| Estimator | Mechanism | Trade‑offs |
|---|---|---|
| TOPK | Computes pmean of per‑device top_k margins |
Cheap, minimal communication, slightly noisy estimate |
| HIST | Bins margins globally into qb_hist_bins, reads quantile from summed histogram |
Smoother estimate, but reduces per‑expert counts due to histogram merging |
The TOPK estimator requires only a small collective operation per layer, making it efficient for large deployments. The HIST estimator calls _bincount_upper_quantile (lines 845–886) to perform a fused bincount over the live computation grid, yielding more stable thresholds at modest additional cost.
Configuring QB Balancing
To use histogram‑based quantile balancing in a Marin model:
from experiments.grug.moe_hero_ep.model import GrugModelConfig, QbEstimator
config = GrugModelConfig(
vocab_size=32000,
hidden_dim=1024,
num_experts=256,
num_experts_per_token=4,
qb_estimator=QbEstimator.HIST, # or QbEstimator.TOPK
qb_hist_bins=1024, # for HIST mode
)
During training, the QB estimator determines β thresholds that guide which experts process each token, while the load‑balancing loss continuously adjusts the router to favor uniform distributions. This dual mechanism prevents both routing instability and expert underutilization.
Integration and Training Flow
The complete MoE expert load balancing with quantile balancing pipeline operates as follows:
- Forward pass: QB router computes margins, estimates quantile via TOPK or HIST, selects experts exceeding threshold β
- Token dispatch: Selected tokens are routed to their assigned experts
- Loss computation:
load_balancing_lossis calculated from actual token fractions and mean routing probabilities - Backward pass: Gradient updates both the router parameters and the quantile estimation statistics
The auxiliary loss is automatically extracted from extras["load_balancing_loss"] and added to the total loss in the training loop.
Key Source Files
| File | Purpose |
|---|---|
experiments/grug/moe_hero_ep/model.py |
GrugModelConfig, QbEstimator enum, load‑balancing loss calculation (lines 439–447), _bincount_upper_quantile histogram helper (lines 845–886) |
lib/levanter/src/levanter/models/moe.py |
Shared MoE utilities including dense_router_delta |
lib/levanter/src/levanter/models/qwen3_moe.py |
HF‑compatible load‑balancing loss for Qwen‑3 MoE reference |
lib/levanter/tests/test_qwen3_moe.py |
Unit tests validating loss computation against HuggingFace implementation |
Summary
- Auxiliary load‑balancing loss:
num_experts * sum(token_fraction * p)penalizes imbalanced routing, computed per layer and summed across the model - Quantile balancing: Dynamically sets expert selection threshold β as the (1 − K/E) quantile of logit margins
- TOPK estimator: Fast, communication‑efficient quantile estimate via per‑device top‑k averaging
- HIST estimator: Smoother global quantile via histogram binning, with configurable
qb_hist_bins - Both mechanisms work together: QB determines immediate routing decisions, load‑balancing loss shapes long‑term router behavior
Frequently Asked Questions
What problem does the load‑balancing loss solve in MoE training?
Without load‑balancing loss, MoE routers tend to collapse to a few dominant experts, leaving most experts untrained and wasting model capacity. The loss explicitly penalizes this by increasing when high‑probability experts also receive high token counts, forcing exploration of all experts.
How do I choose between TOPK and HIST quantile estimators?
Use TOPK when minimizing communication overhead is critical—it's faster and scales well to many devices. Use HIST when routing stability matters more than peak speed, as the histogram produces smoother threshold estimates that reduce variance in expert selection.
Does quantile balancing eliminate the need for load‑balancing loss?
No. Quantile balancing determines which experts are available for selection based on the current logit distribution, but doesn't directly encourage uniform usage. The load‑balancing loss provides the training signal that actually shapes the router to distribute tokens evenly.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →