NeurIPS 2026

Distributionally Robust Mixture-of-Experts Training

Xin TengMuxiao LiHongyi Wen
New York University

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.

+1.42pts
7-task average at 10.3B total params (0.6625 → 0.6767)
−4.3%
excess loss when tokens are forced onto mid-ranked experts
−8.6%
spread (std.) of per-expert loss, mean almost unchanged
+12.9%
expert–domain mutual information: specialization is kept
DRMoET overview. (a) A sparse MoE Transformer with L blocks of attention and an MoE layer produces per-token losses. (b) Inside MoE layer l, a shared expert and top-k routed experts are combined; the router's gate probabilities and the routed experts' output norms form an activation-weighted credit, which gives each expert an attributed loss. (c) Per layer, an EMA of the attributed loss drives a softmax update of the expert weights mu, which puts more mass on high-loss experts; the reweighted DRO loss is backpropagated through the network.
Overview. (a) The backbone is an ordinary sparse MoE Transformer. (b) Inside each MoE layer, DRMoET reads the gate probabilities and the routed experts' output norms to give every expert a detached credit for each token it processed. (c) Credits turn per-token losses into per-expert losses; a per-layer distribution μ moves toward the high-loss experts, and the reweighted loss is what gets backpropagated.
Motivation

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.

Load balancing

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.

DRMoET

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.

Method

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.

\[ \min_{\theta}\; \max_{\mu \in (\Delta_E)^L}\; \sum_{l=1}^{L}\sum_{i=1}^{E} \mu_{l,i}\, R_{l,i}(\theta) \]
Step 1

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.

\[ \tilde c_{l,i}(x_j) = \operatorname{sg}\!\big(p_{l,i}(x_j)\,\lVert h_{l,i}(x_j)\rVert_2\big) \]
Step 2

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).

\[ \ell_{l,i} = \frac{1}{N}\sum_{j} \tilde c_{l,i}(x_j)\,\ell_j\,\mathbf 1[i\in\mathcal T_{j,l}], \qquad \hat\ell_{l,i} \leftarrow \beta\,\hat\ell_{l,i} + (1-\beta)\,\operatorname{sg}(\ell_{l,i}) \]
Step 3

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.

\[ \mu_{l} \leftarrow \operatorname{softmax}\!\big(\mu_{l} + \eta\,\hat\ell_{l}\big), \qquad \mathcal L_{\text{DRO}} = \sum_{l=1}^{L}\sum_{i=1}^{E}\operatorname{sg}(\mu_{l,i})\,\ell_{l,i} \]
Try it

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.

  1. 1EMA of expert loss \(\hat\ell_{l,i}\)
  2. 2Dual update \(\mu_l \leftarrow \operatorname{softmax}(\mu_l + \eta\hat\ell_l)\)
  3. 3Weighted gradient step lowers \(\ell_{l,i}\) in proportion to \(\mu_{l,i}\)
Expert loss \(\ell_{l,i}\)drag to set the startDRMoET uniform
Expert weights \(\mu_l\)change from \(1/E\), % pointsstep 0
Worst-expert loss over stepsDRMoET uniform
After 0 stepsUniformDRMoET
Worst expert––
Mean––
Spread (std.)––
max/min μ1.0001.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.

Results

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.

Analysis

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.

MetricBaselineDRMoETΔ
Worst loss2.69302.6379−2.05%
Mean loss2.37122.3699−0.05%
Std.0.15160.1386−8.58%
CV0.06390.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.

Under distribution shift 60k unseen examples
Loss range1.018 → 0.947−7.0%
Std.0.241 → 0.236−2.1%
Mean loss3.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.

SettingARC-EHellaSwagPIQASciQReCoRD*Average
Credit
FLAME-MoE0.52950.34470.66970.80200.70230.6096
Activation-weighted, η = 10⁻²0.54000.34700.68390.81800.70070.6179
Activation-weighted, η = 10⁻³0.56820.34520.68880.81900.71050.6263
Raw probability, η = 10⁻²0.54550.33840.68390.79200.69380.6107
Raw probability, η = 10⁻³0.54210.34350.68610.79100.70650.6138
Top-12 routing
FLAME-MoE0.54590.35060.69910.79500.70330.6188
DRMoET, η = 10⁻³0.56400.35230.69100.81300.70560.6252
DRMoET, η = 10⁻²0.56190.35150.69910.81100.69980.6247
EMA decay β (η = 10⁻³)
FLAME-MoE0.52950.34470.66970.80200.70230.6096
β = 0.90.55680.34770.69150.79700.70480.6196
β = 0.990.54080.34550.69370.79400.70220.6152
β = 0.999 (default)0.56820.34520.68880.81900.71050.6263
β = 10.54670.34650.68440.78500.70680.6139
Best value per column within each block in bold. Averages are over the five displayed tasks within each block. *ReCoRD is F1; the rest are accuracy.

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.

Cite

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}
}