Distributionally Robust Mixture-of-Experts Training
Sparse routers are never perfect, so a token will sometimes land on an expert that is not its best match. DRMoET treats every routed expert in every layer as a robustness group and shifts training signal toward the experts with the highest attributed loss, so the experts a token falls back on are better trained. The router, architecture and sparse compute stay as they are.
Balanced traffic does not mean competent experts
Imperfect gating, balancing pressure and distribution shift all send tokens to plausible but non-optimal experts. Under standard training those lower-ranked experts get weaker task-aligned updates, so an occasional misroute costs more than it should. Load-balancing losses fix how many tokens each expert sees. They do not check whether an expert does well on the tokens it gets.
Equalize allocation
Penalize skewed routing, e.g. \(L_{\text{aux}} = n\sum_i f_i P_i\), or bias router logits toward under-used experts. Uniform pressure can work against useful specialization.
Improve the worst routing outcomes
Weight each expert's attributed loss by an adversarial distribution, so training targets \(\sum_l \max_i R_{l,i}(\theta)\). Router preferences are left alone.
Experts as layer-wise robustness groups
Group DRO normally needs predefined groups. In an MoE the groups already exist: each expert in each layer handles the tokens the router sends it. DRMoET solves the saddle problem below with three cheap per-step updates, all computed outside autograd.
Credit each routed expert
A token's credit to expert \(i\) is its gate probability times the size of that expert's output, with gradients stopped so the credit never becomes a router or norm objective.
Attribute the token losses
Each expert gets a token-normalized share of the losses on tokens it processed, smoothed over steps with an EMA (\(\beta = 0.999\) by default).
Reweight and backpropagate
A softmax update moves each layer's \(\mu_l\) toward high-loss experts. It is exactly entropic mirror ascent on a regularized robust objective, which gives a \(\tilde{O}(T^{-1/2})\) stationarity guarantee. The reweighted loss replaces the plain average.
The full loop in one layer
Each step runs DRMoET once: the EMA of each expert's loss updates \(\mu_l\), then a gradient step weighted by \(\mu_l\) trains the experts. Upweighted experts get a larger share of the update, so their loss falls faster. Grey ticks show the same experts trained with uniform weights.
- 1EMA of expert loss \(\hat\ell_{l,i}\)
- 2Dual update \(\mu_l \leftarrow \operatorname{softmax}(\mu_l + \eta\hat\ell_l)\)
- 3Weighted gradient step lowers \(\ell_{l,i}\) in proportion to \(\mu_{l,i}\)
| After 0 steps | Uniform | DRMoET |
|---|---|---|
| Worst expert | – | – |
| Mean | – | – |
| Spread (std.) | – | – |
| max/min μ | 1.000 | 1.000 |
A toy simulation to show the mechanism, not the paper's training curves: each expert's loss decays toward a floor (dotted line) at a rate set by its weight \(E\mu_{l,i}\), and the uniform run uses \(\mu_{l,i} = 1/E\) throughout. The μ panel zooms to fit, so small updates stay visible. The paper uses η from 0.0001 to 0.1, where μ stays close to uniform and tilts the gradient gently.
Better downstream accuracy at both scales
From-scratch pretraining on DCLM with the FLAME-MoE recipe. Zero-shot on seven benchmarks (ReCoRD as F1, the rest accuracy). Within each scale, all methods share the architecture, data pipeline and token budget.
Why it works: the fallback experts get stronger
Probes on the best 290M model check the mechanism directly, not just the averages.
Forced mid-k misrouting
Every token's top-k experts are swapped for its middle-k by routing probability, a plausible but non-optimal choice. Loss averaged over six evaluation domains.
Excess loss falls from 4.062 to 3.889: 4.3% less degradation (2.27× → 2.22×).
Per-expert loss at convergence
Measured on a 52M-token validation set. The spread shrinks while the mean barely moves.
| Metric | Baseline | DRMoET | Δ |
|---|---|---|---|
| Worst loss | 2.6930 | 2.6379 | −2.05% |
| Mean loss | 2.3712 | 2.3699 | −0.05% |
| Std. | 0.1516 | 0.1386 | −8.58% |
| CV | 0.0639 | 0.0585 | −8.45% |
Specialization is kept
Six functional domains, averaged over experts; change vs. FLAME-MoE.
- ΔQ competence edge on its own domain
- MI \(I(E;G)\), expert–domain dependence
- ΔS top vs. second domain share
- Max select. top domain share
Mid-tier experts gain most
Experts split into tiers by routing traffic; loss improvement per tier.
The gain lands on the mid tier, where tokens fall back; the busiest experts give up a little.
| Loss range | 1.018 → 0.947 | −7.0% |
| Std. | 0.241 → 0.236 | −2.1% |
| Mean loss | 3.151 → 3.157 | +0.2% |
Negligible extra compute
Per MoE layer and step, with \(N\) tokens, top-\(k\) routing, hidden size \(d\) and FFN width \(d_{\text{ff}}\):
| Expert FFNs (already paid) | \(O(Nk\,d\,d_{\text{ff}})\) |
| Credit norms \(\lVert h_{l,i}\rVert_2\) | \(O(Nk\,d)\) |
| Attribution \(\ell_{l,i}\) | \(O(Nk)\) |
| EMA + softmax for \(\mu_l\) | \(O(E)\) |
The added work is a \(O(1/d_{\text{ff}})\) fraction of the expert compute, the backward pass is still a single pass over a reweighted loss, and the extra state is \(O(LE)\) numbers.
Ablations
290M model. Credit assignment, routing sparsity and EMA smoothing, each against FLAME-MoE trained under the same setting.
| Setting | ARC-E | HellaSwag | PIQA | SciQ | ReCoRD* | Average |
|---|---|---|---|---|---|---|
| Credit | ||||||
| FLAME-MoE | 0.5295 | 0.3447 | 0.6697 | 0.8020 | 0.7023 | 0.6096 |
| Activation-weighted, η = 10⁻² | 0.5400 | 0.3470 | 0.6839 | 0.8180 | 0.7007 | 0.6179 |
| Activation-weighted, η = 10⁻³ | 0.5682 | 0.3452 | 0.6888 | 0.8190 | 0.7105 | 0.6263 |
| Raw probability, η = 10⁻² | 0.5455 | 0.3384 | 0.6839 | 0.7920 | 0.6938 | 0.6107 |
| Raw probability, η = 10⁻³ | 0.5421 | 0.3435 | 0.6861 | 0.7910 | 0.7065 | 0.6138 |
| Top-12 routing | ||||||
| FLAME-MoE | 0.5459 | 0.3506 | 0.6991 | 0.7950 | 0.7033 | 0.6188 |
| DRMoET, η = 10⁻³ | 0.5640 | 0.3523 | 0.6910 | 0.8130 | 0.7056 | 0.6252 |
| DRMoET, η = 10⁻² | 0.5619 | 0.3515 | 0.6991 | 0.8110 | 0.6998 | 0.6247 |
| EMA decay β (η = 10⁻³) | ||||||
| FLAME-MoE | 0.5295 | 0.3447 | 0.6697 | 0.8020 | 0.7023 | 0.6096 |
| β = 0.9 | 0.5568 | 0.3477 | 0.6915 | 0.7970 | 0.7048 | 0.6196 |
| β = 0.99 | 0.5408 | 0.3455 | 0.6937 | 0.7940 | 0.7022 | 0.6152 |
| β = 0.999 (default) | 0.5682 | 0.3452 | 0.6888 | 0.8190 | 0.7105 | 0.6263 |
| β = 1 | 0.5467 | 0.3465 | 0.6844 | 0.7850 | 0.7068 | 0.6139 |
Credit should follow the routed activation. Weighting by gate probability times output norm beats gate probability alone at both η (best average 0.6263 vs 0.6138). This matches the forward pass, where the gate scales the expert's output.
Gains hold under denser routing, but shrink. With top-12 routing DRMoET still improves on FLAME-MoE (0.6252 vs 0.6188). With more experts per token, any single routing choice matters less, so there is less to fix.
Smoothing matters. β = 0.999 gives the best average and is the default. Smaller β follows noisy batch losses more closely and helps individual tasks (β = 0.9 on ARC-E and HellaSwag) but lowers the average.
BibTeX
@inproceedings{teng2026drmoet,
title = {Distributionally Robust Mixture-of-Experts Training},
author = {Teng, Xin and Li, Muxiao and Wen, Hongyi},
booktitle = {Advances in Neural Information Processing Systems},
year = {2026}
}