Paper Summaries
Switch Transformer
Overview
The Switch Transformer makes a language model much larger in parameter count without making it proportionally more expensive to run. It achieves this through sparse activation: each input token activates only a small, selected slice of the model rather than the whole network.
In a standard Transformer, parameters and per-token computation are tightly coupled, so a bigger model means every forward pass costs more. The Switch Transformer decouples them. It replaces the single feed-forward network (FFN) in each Transformer block with a large collection of FFNs called experts, plus a lightweight router that sends each token to just one expert. Adding experts adds parameters, but since each token still visits only one expert, the compute (FLOPs) per token stays roughly constant.
This lets the authors scale to hundreds of billions and even trillions of parameters at a per-token cost comparable to a much smaller dense model, while also improving how fast the model learns per unit of computation. The backbone is T5 (an encoder–decoder Transformer) pretrained on the C4 dataset with a masked span-prediction objective; the Switch modification only changes the FFN sublayers.
Architecture of the Switch Transformer
The architecture is a standard Transformer in which each FFN sublayer is replaced by a Switch layer. A Switch layer has two parts: a set of experts (each its own FFN) and a router that assigns each token to one of them.
The router
Let be the vector representation of a single token arriving at the layer. The router holds a weight matrix that produces routing logits:
has one entry per expert. A softmax turns these logits into a probability distribution over experts:
is the gate value for expert , meaning how strongly the router thinks token belongs to expert . These values are non-negative and sum to 1 across experts.
Switch routing (top-1)
The central idea of the paper is to route each token to a single expert, the one with the highest gate value. The layer output is:
where is the selected expert. Earlier Mixture-of-Experts work argued the router must send each token to at least two experts () to get a useful learning signal. Switch shows works, which brings three benefits: cheaper router computation, smaller per-expert token buffers, and less data communicated between devices in distributed settings.
Note that the expert output is still multiplied by its gate value . This is not just scaling. Because is differentiable, this multiplication is what lets gradients reach the router weights and train the router even under top-1 routing.
How it is trained
The model is trained on the span-prediction objective, but two extra pieces make sparse routing behave well: an auxiliary loss added to the objective, and a mechanical capacity limit applied during the forward pass. This section covers the background it builds on, the loss, and then that capacity mechanism.
Background: Mixture of Experts (MoE)
Switch routing is a simplification of the older Mixture of Experts idea. In standard MoE, the token is routed to the top- experts (those with the highest gate values), and the layer output is the gate-weighted sum of their outputs:
where is the set of selected expert indices. Switch is the special case , i.e. .
Loss function
The full training objective is the primary task loss plus a scaled auxiliary balancing term:
In words: the model is trained mainly to predict the masked spans correctly (), and a second term nudges the router to spread tokens evenly across experts. controls how strongly that balance is enforced (), and is the number of experts. The two terms are unpacked below.
Cross-entropy loss
is the standard span-prediction cross-entropy of the T5 objective, the ordinary language-modeling loss measuring how well the model predicts the masked target tokens. This is the term the model would be trained on even without any expert machinery.
Load balancing loss
Left alone, the router might funnel most tokens to a few favored experts, leaving others idle. The load balancing loss discourages this. For a batch of tokens and experts, define two per-expert quantities.
The fraction of tokens actually dispatched to expert :
where is 1 when the condition holds and 0 otherwise. And the average router probability assigned to expert :
The loss is their scaled dot product:
Here is the hard count of how many tokens went to expert , and is the soft confidence the router expressed toward it. Minimizing their product drives the distribution toward uniform. The minimum is reached when both and are near for every expert. The factor keeps the loss on a consistent scale as the expert count changes. Since is a count (not differentiable) and is differentiable, the gradient flows through the term, giving the router a smooth signal to rebalance while still measuring the actual routing outcome. To understand this clearly see the summed up forward pass below.
Forward pass sum up
Let be the token's vector, and let the router weight matrix be ( = number of experts, = model dimension).
1. Router logits
Each entry is the raw score for expert .
2. Softmax to gate values
is the gate value (probability) for expert .
3. Pick the expert (top-1)
is the index of the chosen expert.
4. Layer output
The chosen expert processes the token, and its output is scaled by that expert's gate value .
Expert capacity
Expert capacity is not a loss. It is a mechanical limit applied during the forward pass, so it sits outside the loss function. Because hardware needs fixed tensor shapes, each expert is given a fixed budget of tokens it can process:
is an absolute count of tokens per expert (e.g. 128 tokens), not a ratio. The first factor is the even share each expert would receive under perfectly uniform routing; the capacity factor () adds slack.
How it is used: each expert's input buffer is pre-sized to hold exactly tokens. As routed tokens fill an expert's buffer in order, any token arriving after the buffer is full overflows and is dropped, receiving no expert output and passing forward only through the residual connection. Experts with fewer than tokens leave the remaining slots as padding but still compute over the full block, which is the wasted-compute cost of a high capacity factor. This mechanism runs on every forward pass, so it is active during both training and inference; the load balancing loss exists precisely to keep buffers from overflowing.
Training stability and fine-tuning
Large sparse models are prone to instability, so the paper adds several techniques.
Selective precision
Low-precision training (bfloat16) is efficient, but the router's softmax is numerically sensitive because exponentials amplify small perturbations. The fix is selective precision: the router's internal computation is done in float32 and cast back to bfloat16 afterward. Since this float32 region is local to the router and not communicated between devices, stability is gained without the communication cost of full float32.
Smaller initialization
Reducing the weight initialization scale improved stability. The initialization scale factor is cut by a factor of 10 (from to ) when drawing from a truncated normal distribution, lowering the risk of exploding activations and gradients early in training.
Expert dropout
When fine-tuning on smaller downstream tasks, overfitting is a risk. The paper applies a higher dropout rate inside the experts than elsewhere (around 0.4 at the experts vs. 0.1 elsewhere). Because experts hold most of the parameters, concentrating regularization there improves fine-tuning without over-regularizing the shared parts of the network.
Scaling and parallelism
To train models this large, three forms of parallelism are combined. Data parallelism places different batches on different devices; model parallelism splits individual weight tensors across devices; and expert parallelism places different experts on different devices, so a token routed to a given expert is sent to that expert's device. Switch fits naturally with expert parallelism, and the three can be combined so both the number of experts and the size of each expert grow with the available hardware.
Empirically, at fixed compute per token the Switch Transformer learns substantially faster and reaches better quality than a dense T5 baseline, up to about 7× pretraining speedup to a fixed quality. Scaling the expert count consistently helps, culminating in Switch-C, a model of roughly 1.6 trillion parameters using 2048 experts.
Distillation to dense models
A large sparse model is powerful but cumbersome to deploy. A trained Switch model can be distilled back into a much smaller dense model, transferring roughly 30% of the sparse model's quality gains into a compact student that is far cheaper to serve, a path from large-scale sparse training to practical dense deployment.
Summary
In simple terms: during training you train the Switch Transformer architecture with cross-entropy loss plus load balancing loss, while considering expert capacity; and during inference you also consider expert capacity.
Appendix (Forward pass with shapes)
Setup and dimensions
Input to the Switch layer: . I'll use for batch (you wrote , but I'll reserve for the number of experts to avoid a clash).
- = batch size
- = sequence length
- = model dimension
- = number of experts
- = expert hidden dimension (the FFN's inner width)
A useful move: routing happens per token, and there are tokens total. So it's common to flatten:
Now every token is just a row of dimension , and there are of them.
Forward pass with shapes
Router weight matrix
Maps a -dim token to logits (one per expert).
1. Router logits, for all tokens at once:
So , each of the tokens gets logits. For a single token, .
2. Softmax to gate values (softmax across the axis):
. Row is the gate distribution over experts for token . For a single token, , a vector of probabilities.
3. Pick the expert (top-1), argmax across the axis:
, one integer index per token. And (gathering each token's chosen gate value) is:
a single scalar per token. For one token, is just a scalar .
4. Experts
Each expert is an FFN with two linear layers:
For one token routed to expert :
So , same dimension as the input token. Stacked over all experts, the full expert parameters are and , but each token only touches the one slice indexed by .
5. Layer output, scale the chosen expert's output by the chosen gate scalar:
For all tokens: , which reshapes back to , the same shape as the input.