Model Efficiency & Scaling2022advanced13 min read

Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity

المحوِّلات التبديلية: التوسّع إلى تريليون معامل بتناثر بسيط وفعّال

Fedus, W. · Zoph, B. · Shazeer, N. — JMLR

The problem

By 2021, scaling neural language models was proving enormously effective — bigger models performed better. But making models bigger meant proportionally more computation for every single . A model with 10× more parameters needed roughly 10× more FLOPs per token. (MoE) offered a theoretical escape: increase parameters without increasing per-token compute. But in practice, MoE models were plagued by instability, implementation complexity, and high communication costs. The top-k routing strategy required sending each token to multiple experts, adding overhead. No one had managed to scale MoE simply and stably to the trillion-parameter regime.

The contribution

The Switch simplifies MoE routing to its extreme: each token goes to exactly one expert (top-1), not two or more. This halves the expert compute, reduces communication costs, and simplifies implementation — while maintaining or improving quality. The paper introduces training stabilization techniques (selective float32 precision for the router, reduced initialization scale), an strategy for , and a differentiable load-balancing auxiliary loss. Built on top of T5, Switch Transformers achieve up to 7× speedup at the same compute budget, scale to 1.6 trillion parameters, improve across all 101 languages in multilingual settings, and can be distilled back into small dense models preserving ~30% of quality gains.

The impact

Switch Transformers proved that sparse expert models are practical, not just theoretical. By demonstrating that top-1 routing works — contradicting the prevailing belief that top-k>1 was necessary — the paper opened the door to simpler, larger MoE architectures. This directly influenced Mixtral (Mistral AI), DeepSeek-V3, and other modern MoE models that now power production systems. The scaling insight — that parameter count matters independently of FLOPs — became a guiding principle for the field. The paper also showed that sparse models can be compressed via distillation for practical deployment, bridging research scale with engineering reality.

In a standard Transformer, every token passes through the same feed-forward network — like a single doctor examining every patient in the hospital, no matter the specialty needed. The doctor is brilliant, but there's only one of them.

Mixture of Experts says: hire 128 doctors. But previous MoE systems sent every patient to two or more specialists for a second opinion, creating scheduling nightmares and doubling the workload.

The Switch Transformer's insight is ruthlessly simple: one patient, one doctor. A smart receptionist (the router) glances at symptoms and makes a single assignment. The hospital has 128× the expertise, but each visit costs the same as before.

The problem: scaling means paying more per token

The dominant recipe for improving language models in 2020 was straightforward: make the model bigger. GPT-3 showed that scaling to 175 billion parameters produced remarkable few-shot abilities. But every new parameter participated in every computation — 10× more parameters meant 10× more FLOPs per token. This made training prohibitively expensive.

Kaplan et al. (2020) discovered power-law scaling relationships between model size, data, compute, and loss — but all three axes were coupled. Could we decouple parameter count from compute? Mixture of Experts (MoE) promised exactly that: a model with billions of parameters where each token only activates a small subset. The parameters grow, but the FLOPs per token stay constant.

Yet MoE had three stubborn problems. First, routing tokens to the right expert was complex — the standard approach sent each token to at least two experts. Second, training was unstable, especially at lower precision formats like bfloat16. Third, all-to-all communication between devices for token dispatching was expensive. The Switch Transformer addresses all three.

The core idea: route each token to exactly one expert

In a standard Transformer, each layer contains a sublayer followed by a feed-forward network (). The Switch Transformer replaces the FFN with a Switch FFN layer: instead of one FFN applied to every token, there are N independent FFN experts, and a lightweight router picks exactly one expert for each token.

The previous state-of-the-art, Shazeer et al. (2017), had conjectured that routing to k > 1 experts was necessary for meaningful flow to the router. The Switch Transformer challenges this directly: k = 1 works. And it brings three immediate benefits: the router computation is reduced (only pick one, not rank two), each expert's can be halved (tokens go to one expert, not two), and communication across devices is simpler.

The router itself is a single learned linear layer. It takes each token's representation x, multiplies by a weight matrix W_r to produce logits over N experts, applies to get probabilities, and sends the token to the highest-probability expert. The output is the expert's computation scaled by the router's gate probability.

Open in Lab
Toggle between top-1 (Switch) and top-2 (MoE) routing. Click tokens to see router probabilities — notice Switch halves the compute.
The demo wakes as you arrive…
pi(x)=eh(x)i∑j=1Neh(x)j,h(x)=Wr⋅xp_i(x) = \frac{e^{h(x)_i}}{\sum_{j=1}^{N} e^{h(x)_j}}, \quad h(x) = W_r \cdot x
Router probability — softmax over expert logits — The router computes a probability distribution over N experts for each token x. The token is sent to the expert with the highest probability. The gate value p_i(x) scales the expert's output, preserving differentiability.
Open in Lab
Compare the dense Transformer block (left) with the Switch Transformer block (right). Click any component to learn its role.
The demo wakes as you arrive…

Expert capacity: when experts overflow

In distributed training, tensor shapes must be statically declared. Each expert is allocated a fixed batch size — its capacity — computed as:

= (tokens per batch / number of experts) ×

If the capacity factor is 1.0, experts have exactly enough room for a perfectly balanced distribution. In practice, routing is never perfectly uniform, so some experts overflow. Overflowing tokens are dropped — they skip the expert layer entirely and pass through via the , receiving no expert processing.

A larger capacity factor (e.g. 1.5) adds buffer slots, reducing drops but wasting compute on empty padding. A smaller capacity factor (e.g. 1.0) is more efficient but risks more token drops. The paper found that capacity factor 1.0 to 1.25 works best in practice, especially in the large-scale regime where memory is scarce.

Expert Capacity=tokens per batchnum experts×capacity factor\text{Expert Capacity} = \frac{\text{tokens per batch}}{\text{num experts}} \times \text{capacity factor}
Expert capacity formula — Each expert can process at most this many tokens. Tokens exceeding this limit are dropped and skip to the next layer via the residual connection.
Open in Lab
Adjust the capacity factor to see the trade-off: too small drops tokens, too large wastes compute.
The demo wakes as you arrive…

Load balancing: making sure every expert earns its keep

Without intervention, routers tend to converge to sending most tokens to just a few experts while the rest sit idle — a classic rich-get-richer problem. If Expert 1 receives more tokens early on, it gets better gradients, becomes more useful, and attracts even more tokens. This imbalance wastes parameters and causes token overflow.

The fix is an auxiliary load-balancing loss added to the training objective. For each Switch layer, we compute two vectors over the N experts: f (the actual fraction of tokens dispatched to each expert) and P (the mean router probability assigned to each expert). The auxiliary loss is their scaled dot product:

If routing is perfectly uniform, both vectors are (1/N, ..., 1/N) and the loss is minimized. Any deviation increases the loss. The key insight is that while f is not differentiable (it's based on argmax decisions), P is differentiable through the softmax, so the loss can still push the router toward balance. The coefficient α = 10⁻² was found large enough to ensure while small enough not to interfere with the primary language modeling objective.

Laux=α⋅N⋅∑i=1Nfi⋅Pi\mathcal{L}_{\text{aux}} = \alpha \cdot N \cdot \sum_{i=1}^{N} f_i \cdot P_i
Auxiliary load-balancing loss — f_i is the fraction of tokens dispatched to expert i (not differentiable). P_i is the mean router probability for expert i (differentiable through softmax). Their dot product is minimized when routing is uniform. α = 10⁻² and N scales the loss.
Open in Lab
Increase α to see how the auxiliary loss pushes token distribution toward uniform.
The demo wakes as you arrive…

Training stability: taming the instabilities

Sparse expert models face unique training challenges. The hard routing decisions create discontinuities, and low-precision formats amplify numerical errors in the softmax computation. The paper introduces two stabilization techniques:

Selective precision. Instead of training entirely in float32 (stable but slow) or entirely in bfloat16 (fast but diverges), the Switch Transformer casts only the router function to float32. The dispatch and combine tensors are recast back to bfloat16 before all-to-all communication, so the expensive cross-device transfers stay in low precision. Result: bfloat16 speed with float32 stability.

Reduced initialization scale. Standard Transformer initialization uses a scale factor s = 1.0. Switch Transformers reduce this to s = 0.1 — a 10× smaller initialization. This dramatically improved quality and reduced variance across runs. The same scheme worked from 223M parameters all the way to 1 trillion+ parameters.

Open in Lab
Compare three precision modes: full float32 is stable but slow, full bfloat16 diverges, selective precision gets the best of both.
The demo wakes as you arrive…

Scaling: more experts, faster learning

The most efficient dimension for scaling the Switch Transformer is the number of experts. Increasing experts keeps FLOPs per token approximately constant — the router just computes a distribution over more experts, an O(d_model × num_experts) operation that is negligible compared to the FFN computation.

The paper systematically scaled from 2 to 256 experts, observing consistent improvements in both step-efficiency and wall-clock time. Key findings:

  • A Switch-Base with 64 experts achieves the same quality as T5-Base 7.5× faster in terms of training steps, and 7× faster in wall-clock time.

  • Even compared to T5-Large, which uses 3.5× more FLOPs per token, Switch-Base with 64 experts is 2.5× faster on a wall-clock basis.

  • Models with as few as 2 experts already show meaningful improvements over dense baselines, making the approach useful even on limited hardware.

The scaling also holds in the trillion-parameter regime. Switch-C (1.6 trillion parameters, 2048 experts) achieves a 4× speedup over T5-XXL with stable training throughout.

Open in Lab
Toggle models to compare pre-training speed. All use the same FLOPs — more experts means faster convergence.
The demo wakes as you arrive…

Downstream results: pre-training gains transfer

Upstream pre-training improvements are only valuable if they transfer to real tasks. The paper validates this across a diverse set of benchmarks:

  • SuperGLUE: Switch-Base improves 4.4 points over T5-Base; Switch-Large improves 2 points over T5-Large. These are FLOP-matched comparisons — same compute, better results.

  • SQuAD: Switch-Base achieves 87.2 (vs 85.5 for T5-Base). Switch-Large reaches 88.6.

  • Winogrande: Switch-Large jumps to 83.0 (vs 79.1 for T5-Large), showing gains on commonsense reasoning, not just pattern matching.

  • Closed-book QA: Large improvements on Trivia QA (30.7 vs 24.5 for T5-Base) suggest that sparse models store more factual knowledge in their expanded parameter space.

  • Multilingual: Across all 101 languages in mC4, the Switch model improves over the mT5-Base baseline. 91% of languages see at least 4× speedup. This is especially significant for low-resource languages that benefit from the shared capacity.

Distillation: compressing trillions into millions

A trillion-parameter model is impractical to deploy. The paper shows that sparse teachers can be distilled into small, dense student models using two key techniques:

Non-expert : since the Switch model is FLOP-matched to the dense baseline, all non-expert layers (self-attention, layer norms, embeddings) have the same dimensions. These trained weights are used to initialize the dense student, giving it a head start.

Mixed loss: instead of training the student purely on ground truth, the loss is a mixture of 75% hard labels (ground truth) and 25% soft labels (teacher's output probabilities). The soft labels encode the teacher's "" — its distribution over all possible answers.

With both techniques combined, a 14.7B parameter Switch model can be compressed 99% into a 223M dense student while preserving 28% of the quality gain. At 82% compression (1.1B → 223M), 37% of quality is preserved. These are practical compression rates for deployment.

Open in Lab
Select a teacher model size to see how much quality survives compression to 223M params.
The demo wakes as you arrive…

The idea in code

Switch routing — core logic for top-1 expert selectionpython

Simplified to show the idea — not the real implementation.

import numpy as np

def switch_router(token_reprs, expert_weights, num_experts):
    """Route each token to its best single expert (top-1).

    Args:
        token_reprs: (batch, d_model) — token representations
        expert_weights: (d_model, num_experts) — router weight matrix
        num_experts: number of available experts
    Returns:
        expert_indices: (batch,) — which expert each token goes to
        gate_values: (batch,) — softmax probability for scaling output
    """
    # Compute router logits
    logits = token_reprs @ expert_weights          # (batch, num_experts)

    # Softmax to get probabilities (cast to float32 for stability!)
    logits_f32 = logits.astype(np.float32)
    probs = np.exp(logits_f32 - logits_f32.max(axis=-1, keepdims=True))
    probs /= probs.sum(axis=-1, keepdims=True)

    # Top-1 selection — the "Switch" in Switch Transformer
    expert_indices = np.argmax(probs, axis=-1)     # (batch,)
    gate_values = probs[np.arange(len(probs)), expert_indices]

    return expert_indices, gate_values

def load_balance_loss(probs, expert_indices, num_experts, alpha=0.01):
    """Auxiliary loss encouraging uniform expert usage."""
    # f_i: fraction of tokens dispatched to expert i
    f = np.zeros(num_experts)
    for idx in expert_indices:
        f[idx] += 1
    f /= len(expert_indices)

    # P_i: mean router probability for expert i
    P = probs.mean(axis=0)

    # Loss = α * N * Σ(f_i * P_i), minimized when both are uniform
    return alpha * num_experts * np.sum(f * P)

Parallelism: data, model, and expert dimensions

Scaling to a trillion parameters requires combining three parallelism strategies:

splits the batch across cores. Each core has the full model but sees different data. Communication only happens when gradients are aggregated at the end of each step.

splits the model's weight matrices across cores. Each core holds part of the weights but processes the full batch. This requires all-reduce communication in every forward and backward pass.

is unique to MoE models. Each core holds a different expert. Tokens are dispatched to the correct expert via all-to-all communication. This naturally maps to the Switch Transformer: with N cores and N experts, each core owns one expert.

The paper's two flagship models use different combinations: Switch-C (1.6T params) uses only expert + data parallelism with 2048 experts, keeping the per-expert model small. Switch-XXL (395B params) uses all three, with larger per-expert dimensions but higher communication overhead. Switch-C exhibited no training instability; Switch-XXL occasionally did, highlighting the trade-off between FLOPs per token and stability.

From Switch to modern MoE

  1. 1991

    Mixture of Experts (Jacobs et al.)

    The original MoE framework: multiple expert networks with a gating mechanism to combine their outputs. Foundational idea, but limited to small scale.

  2. 2017

    Sparsely-Gated MoE (Shazeer et al.)

    First modern MoE in deep learning: MoE layers between LSTMs with top-k routing. SOTA in translation and language modeling. Proved MoE works at scale.

  3. 2020

    GShard (Lepikhin et al.)

    MoE integrated into the Transformer for massively multilingual translation (100 languages). Extended XLA compiler for automatic sharding.

  4. 2022

    Switch Transformer (Fedus et al.)

    Simplifies MoE to top-1 routing. First stable training at trillion-parameter scale. 7× speedup over T5. Distillation to deployable sizes.

  5. 2024

    Mixtral (Mistral AI)

    Open-weight MoE model using 8 experts with top-2 routing per layer. Competitive with much larger dense models. Direct descendant of Switch ideas.

  6. 2024

    DeepSeek-V3

    Large-scale MoE model with innovative auxiliary-loss-free load balancing and multi-head latent attention. Pushes MoE into production frontier models.

CitationFedus, Zoph, Shazeer. Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity. JMLR, 2022.

Terms in this paper