# How Marin Implements Mixture-of-Experts Quantile Balancing: A Technical Deep Dive

> Discover how Marin implements mixture-of-experts quantile balancing with QB routers, dynamic thresholds, and expert load estimation to prevent overload. A technical deep dive.

- Repository: [The Marin Project/marin](https://github.com/marin-community/marin)
- Tags: deep-dive
- Published: 2026-08-28

---

**Marin implements mixture-of-experts quantile balancing through a QB router that computes per-token margins against a dynamic threshold α, then estimates expert load distributions using either a top-k sampling or histogram-based method to iteratively adjust router biases and prevent expert overload.**

The marin-community/marin repository introduces a sophisticated approach to mixture-of-experts quantile balancing that prevents load imbalance in massive MoE architectures. By combining top-k routing with adaptive threshold margins and statistical quantile estimation, Marin's implementation ensures efficient token distribution across hundreds of experts. This article examines the core mechanisms, from margin calculation to bias application, based on the actual source implementation in the Grug experimental framework.

## Top-K Plus QB Routing Architecture

At the heart of Marin's approach is the **Top-K + QB Routing** strategy implemented in `MoEMLP.__call__` within [`experiments/grug/moe_hero_ep/model.py`](https://github.com/marin-community/marin/blob/main/experiments/grug/moe_hero_ep/model.py). For every input token, the router computes raw logits (`router_logits`) and applies a learnable bias to produce biased logits.

The routing process follows these steps:

1. **Threshold Selection** – The router selects the top-`K + 1` experts using the biased logits (`router_logits + router_bias`). The `(K + 1)`-st logit defines the **QB threshold** `α` (alpha), which serves as a dynamic margin boundary.

2. **Margin Calculation** – The top-`K` experts process the token, while the difference `s − α` (where `s` represents the logit scores of selected experts) forms the **margin** used for quantile estimation.

According to the source code at lines 974-979, this logic extracts `qb_alpha` and prepares the margins for downstream quantile computation:

```python

# Simplified excerpt from MoEMLP.__call__

# router_logits: [batch, seq, num_experts]

# router_bias: [num_experts]

biased_logits = router_logits + router_bias
top_k_plus_one_probs, top_k_plus_one_indices = jax.lax.top_k(biased_logits, self.k + 1)
qb_alpha = top_k_plus_one_probs[..., -1:]  # (K+1)th expert defines threshold

margins = selected_probs - qb_alpha  # s - α

```

## Quantile Estimation Modes

Marin supports two distinct strategies for estimating the per-expert quantile `β` via the `QbEstimator` enum defined in the configuration. The choice between these estimators allows trading off computational cost against statistical precision.

### TOPK Estimator

The **TOPK** mode (`QbEstimator.TOPK`) computes a per-shard quantile by extracting the top-k values of the margins on each device, then averaging these values across the distributed setup. This approach is computationally efficient but produces noisier estimates since it only samples local data.

When `cfg.qb_estimator == QbEstimator.TOPK`, the code path falls back to a simple per-shard top-k computation inside `MoEMLP.__call__` (lines 1005-1020), storing the result in `router_stats["qb_beta_local"]` rather than a globally synchronized quantile.

### HIST Estimator

The **HIST** mode (`QbEstimator.HIST`) provides a more robust estimate by building a histogram of margins across a global `[min, max]` grid. This method reads the `(1 − K/E)`-quantile from the distribution, where `E` represents the total number of experts.

The heavy lifting occurs in `_qb_beta_hist` (lines 1084-1092), which first gathers the global margin range (`pmin/pmax`) across all devices, then delegates to `_bincount_upper_quantile` for the actual quantile extraction.

## Histogram-Based Quantile Calculation

The histogram implementation in `_bincount_upper_quantile` (lines 845-872) represents the most statistically rigorous component of Marin's quantile balancing system. This function operates through the following pipeline:

1. **Binning** – Margins are distributed into `n_bins` across each expert's `[lo, hi]` interval.

2. **Global Aggregation** – A single fused `jnp.bincount` operation creates a local histogram over the flattened `(expert × bin)` index. This histogram is then summed across all devices using `jax.lax.psum` to achieve a global view of the margin distribution.

3. **Quantile Extraction** – The cumulative counts from the top of the histogram are examined to locate the bin where the **target rank** (`target_rank = tokens × K/E`) falls. Linear interpolation inside that bin yields the final quantile value `β`.

```python

# Conceptual flow inside _bincount_upper_quantile

local_hist = jnp.bincount(flattened_indices, length=num_experts * n_bins)
global_hist = jax.lax.psum(local_hist, axis_name="devices")

# Find (1-K/E) upper quantile via cumulative counts from top

cumulative = jnp.cumsum(global_hist[::-1])[::-1]
target_rank = total_tokens * k / num_experts

# ... linear interpolation to find beta

```

## Router Statistics and Bias Application

After computing the per-layer quantile `β`, Marin stores this value in `router_stats["qb_beta"]` (or `qb_beta_local` for the TOPK path). The training loop then applies this quantile as a negative bias to the router logits, creating a feedback mechanism that adjusts the selection threshold for the next forward pass.

In [`experiments/june_tpu_67b_a2b/moe/train.py`](https://github.com/marin-community/marin/blob/main/experiments/june_tpu_67b_a2b/moe/train.py) (lines 332-334), the `_apply_qb_betas` utility encapsulates this logic:

```python
def _apply_qb_betas(model, pending_qb_betas):
    # Updates router_bias = -beta for the next step

    model = model.replace(
        router_bias=model.router_bias - pending_qb_betas
    )
    return model

```

This mechanism ensures that experts receiving excessive load (indicated by high margins) will have their effective thresholds raised, naturally diverting tokens toward underutilized experts in subsequent iterations.

## Configuring Quantile Balancing in Marin

To enable mixture-of-experts quantile balancing in your Marin model, configure the `GrugModelConfig` with the appropriate estimator settings:

```python
from levanter.grug.moe import QbEstimator, GrugModelConfig

cfg = GrugModelConfig(
    vocab_size=32000,
    num_experts=256,
    num_experts_per_token=4,
    qb_estimator=QbEstimator.HIST,      # Choose histogram-based estimation

    qb_hist_bins=1024,                  # Resolution for margin histogram

)

model = cfg.build(Vocab=Axis("vocab", cfg.vocab_size), key=jax.random.PRNGKey(0))

```

During training, the quantile balancing loop integrates seamlessly with the forward pass:

```python

# Inside a training step (simplified)

router_stats = moe_layer(x)               # Returns (output, router_stats)

beta = router_stats["qb_beta"]            # Extract per-expert quantile

model = apply_qb_betas(model, beta)       # Update router_bias = -beta

```

## Summary

- **Dynamic Thresholding**: Marin's QB router uses the `(K+1)`-st expert logit as a dynamic threshold `α`, computing margins as `s − α` to measure relative expert affinity.
- **Dual Estimation Strategies**: The system supports `TOPK` for fast, local estimation and `HIST` for precise, global quantile calculation via distributed histograms.
- **Statistical Rigor**: The `_bincount_upper_quantile` function computes the `(1 − K/E)`-upper quantile using `jnp.bincount` and `jax.lax.psum`, with linear interpolation for sub-bin precision.
- **Adaptive Feedback**: Computed quantiles `β` are applied as negative biases (`router_bias = -β`) to iteratively balance expert loads across training steps.
- **Implementation Location**: Core logic resides in [`experiments/grug/moe_hero_ep/model.py`](https://github.com/marin-community/marin/blob/main/experiments/grug/moe_hero_ep/model.py) with training integration in [`experiments/june_tpu_67b_a2b/moe/train.py`](https://github.com/marin-community/marin/blob/main/experiments/june_tpu_67b_a2b/moe/train.py).

## Frequently Asked Questions

### How does the QB threshold α get calculated?

The threshold `α` is derived from the `(K+1)`-th highest biased logit during the top-k selection process. Specifically, `MoEMLP.__call__` selects the top-`K + 1` experts and uses the lowest of these values as the alpha margin against which all selected expert scores are compared. This value represents the "rejected" expert's affinity score, establishing a natural cutoff point for quantile calculation.

### What is the difference between TOPK and HIST estimators?

The **TOPK** estimator calculates quantiles using only the highest margin values from each device's local shard, making it faster but statistically noisier due to limited sampling. The **HIST** estimator constructs a global histogram of all margins across devices using `jax.lax.psum`, then computes the precise `(1 − K/E)`-quantile from the full distribution. HIST provides smoother, more stable balancing at the cost of increased communication overhead.

### Why is the router bias updated as negative beta?

The bias is set to `‑β` because `β` represents the estimated upper quantile of the margin distribution. By subtracting this value from the raw logits in the next forward pass, the router effectively raises the selection threshold for overloaded experts (those with margins near or above `β`). This negative feedback shifts the decision boundary downward, redirecting tokens toward experts with lower historical loads and maintaining equilibrium.

### Where is the quantile balancing logic located in the codebase?

The primary routing logic resides in [`experiments/grug/moe_hero_ep/model.py`](https://github.com/marin-community/marin/blob/main/experiments/grug/moe_hero_ep/model.py), specifically within the `MoEMLP.__call__` method (lines 974-979 for threshold extraction, 1005-1020 for TOPK estimation). Histogram-based quantile calculation is implemented in `_bincount_upper_quantile` (lines 845-872) and `_qb_beta_hist` (lines 1084-1092). The training integration that applies computed biases back to the model is found in `_apply_qb_betas` within [`experiments/june_tpu_67b_a2b/moe/train.py`](https://github.com/marin-community/marin/blob/main/experiments/june_tpu_67b_a2b/moe/train.py) (lines 332-334).