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 tokens
  • num_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:

  1. Forward pass: QB router computes margins, estimates quantile via TOPK or HIST, selects experts exceeding threshold β
  2. Token dispatch: Selected tokens are routed to their assigned experts
  3. Loss computation: load_balancing_loss is calculated from actual token fractions and mean routing probabilities
  4. 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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →