# How the Switch Transformer Routing Mechanism Works: Token-to-Expert Selection Explained

> Learn how the Switch Transformer routing mechanism assigns tokens to experts with softmax, enforces capacity limits, and scales output by confidence for efficient sparse training.

- Repository: [labml.ai/annotated_deep_learning_paper_implementations](https://github.com/labmlai/annotated_deep_learning_paper_implementations)
- Tags: deep-dive
- Published: 2026-03-04

---

**The Switch Transformer routing mechanism assigns each input token to a single feed-forward expert using a learned linear router with softmax probabilities, enforces per-expert capacity limits with optional token dropping, and scales the final output by routing confidence to maintain gradient flow during sparse mixture-of-experts training.**

The Switch Transformer implements a sparse mixture-of-experts (MoE) layer that routes tokens individually rather than processing all tokens through every expert. According to the labmlai/annotated_deep_learning_paper_implementations repository, this routing mechanism is implemented in [`labml_nn/transformers/switch/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/switch/__init__.py) and combines top-1 expert selection with capacity management to enable scalable training.

## Router Construction and Probability Computation

The routing process begins in the `SwitchFeedForward` class, which creates a learned projection from the model dimension to the number of experts. In the constructor (lines 79‑80), a linear layer maps token embeddings to expert logits, followed by a softmax to produce a probability distribution:

```python

# SwitchFeedForward.__init__

self.switch = nn.Linear(d_model, n_experts)
self.softmax = nn.Softmax(dim=-1)

```

During the forward pass, the input tensor is flattened to `[tokens, d_model]` and passed through the router. The softmax operation yields `route_prob`, a probability vector indicating the affinity between each token and every expert:

```python
route_prob = self.softmax(self.switch(x))

```

This probability distribution determines which expert will process each token in the subsequent selection step.

## Top-1 Expert Selection and Token Grouping

The Switch Transformer uses **top-1 routing**, selecting only the expert with the highest probability for each token. The `torch.max` operation extracts both the winning expert index and the corresponding confidence score (line 100):

```python
route_prob_max, routes = torch.max(route_prob, dim=-1)

```

The resulting `routes` tensor contains integer indices indicating the assigned expert for every token. To prepare for parallel expert execution, tokens are grouped by their assigned expert using boolean indexing:

```python
indexes_list = [torch.eq(routes, i).nonzero(as_tuple=True)[0] for i in range(self.n_experts)]

```

This creates a list of index tensors, where `indexes_list[i]` contains the positions of all tokens routed to expert *i*.

## Capacity Management and Token Dropping

To prevent computational overload on individual experts, the mechanism enforces a **capacity factor** that limits how many tokens each expert can process. The capacity is computed as (lines 108‑112):

```python
capacity = int(self.capacity_factor * len(x) / self.n_experts)

```

If the optional `drop_tokens` flag is enabled and an expert receives more tokens than its capacity, the surplus is randomly shuffled and truncated. Dropped tokens bypass expert computation entirely, preserving their original values and gradients (lines 119‑130):

```python
if self.drop_tokens:
    for i in range(self.n_experts):
        if len(indexes_list[i]) > capacity:
            indexes_list[i] = indexes_list[i][torch.randperm(len(indexes_list[i]))]
            dropped.append(indexes_list[i][capacity:])
            indexes_list[i] = indexes_list[i][:capacity]

```

This shuffling ensures fairness when selecting which tokens to drop, and the `dropped` list tracks these tokens for later reintegration into the output.

## Expert Execution and Output Scaling

Each expert—typically a standard feed-forward network—processes only its assigned token subset. The implementation uses list comprehensions to parallelize expert computation across the batch, then writes results back into the appropriate output positions (lines 132‑138):

```python
expert_output = [self.experts[i](x[indexes_list[i], :]) for i in range(self.n_experts)]
for i in range(self.n_experts):
    final_output[indexes_list[i], :] = expert_output[i]

```

After gathering expert outputs, the mechanism applies **routing probability scaling** to weight each token by its selection confidence. If `is_scale_prob` is true, outputs multiply directly by `route_prob_max`. Otherwise, a straight-through estimator preserves gradient flow by dividing the probability by its detached value (lines 144, 148):

```python
if self.is_scale_prob:
    final_output = final_output * route_prob_max.view(-1, 1)
else:
    final_output = final_output * (route_prob_max / route_prob_max.detach()).view(-1, 1)

```

This scaling ensures that experts receiving high-confidence assignments have greater influence on the final representation while maintaining backpropagation through the discrete routing decision.

## Integration into the Transformer Stack

The `SwitchTransformerLayer` integrates the routing mechanism into the standard transformer pipeline by calling `SwitchFeedForward` within its forward method. The layer captures and returns routing statistics—including expert counts, routing probabilities, and drop rates—for load-balancing loss computation (line 107):

```python
ff, counts, route_prob, n_dropped, route_prob_max = self.feed_forward(z)

```

These statistics, tracked across multiple layers in `SwitchTransformer`, enable auxiliary losses that encourage balanced expert utilization during training. The full implementation, including training loops and load-balancing logic, is available in [`labml_nn/transformers/switch/experiment.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/switch/experiment.py).

## Practical Implementation Example

The following example demonstrates how to construct a Switch Transformer manually using the labml library components:

```python
import torch
import torch.nn as nn
from labml_nn.transformers.switch import SwitchTransformer, SwitchTransformerLayer, SwitchFeedForward
from labml_nn.transformers import MultiHeadAttention
from labml_nn.transformers.feed_forward import FeedForward

d_model = 128
heads = 4
d_ff = 256
n_experts = 4
capacity_factor = 1.2
drop_tokens = True
is_scale_prob = True

# Construct a single SwitchTransformerLayer

layer = SwitchTransformerLayer(
    d_model=d_model,
    attn=MultiHeadAttention(heads, d_model, dropout=0.0),
    feed_forward=SwitchFeedForward(
        capacity_factor=capacity_factor,
        drop_tokens=drop_tokens,
        is_scale_prob=is_scale_prob,
        n_experts=n_experts,
        expert=FeedForward(d_model, d_ff, dropout=0.0),
        d_model=d_model),
    dropout_prob=0.0)

# Stack 6 layers to form the complete model

model = SwitchTransformer(layer, n_layers=6)

# Forward pass with shape (seq_len, batch, d_model)

x = torch.randn(64, 32, d_model)
mask = torch.ones(64, 64).bool()
out, counts, route_prob, n_dropped, route_prob_max = model(x, mask)

```

For a complete training example with load-balancing loss, run the predefined experiment:

```python
from labml_nn.transformers.switch.experiment import main

if __name__ == '__main__':
    main()

```

## Summary

- **Learned Router**: A linear layer maps tokens to expert logits, applying softmax to generate routing probabilities in [`labml_nn/transformers/switch/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/switch/__init__.py) (lines 79‑96).
- **Top-1 Selection**: Each token is assigned to the single expert with highest probability using `torch.max`, with indices grouped for batched processing (line 100).
- **Capacity Enforcement**: A configurable `capacity_factor` limits tokens per expert; surplus tokens are randomly dropped when `drop_tokens=True` (lines 112‑130).
- **Gradient Scaling**: Outputs scale by routing probability or use a straight-through estimator to maintain backpropagation through discrete routing decisions (lines 144‑148).
- **Statistics Tracking**: The `SwitchTransformerLayer` returns routing metadata—`counts`, `route_prob`, `n_dropped`, and `route_prob_max`—to support load-balancing auxiliary losses during training (line 107).

## Frequently Asked Questions

### What is the capacity factor in the Switch Transformer routing mechanism?

The **capacity factor** is a hyperparameter that determines how many tokens each expert can process relative to the average load. It is calculated as `int(capacity_factor * total_tokens / n_experts)` (line 112). A value of 1.0 allows exactly the average number of tokens per expert, while values greater than 1.0 provide buffer capacity to handle imbalanced routing distributions. If an expert receives more tokens than its capacity and `drop_tokens` is enabled, the excess tokens are randomly discarded to maintain computational constraints.

### How does the Switch Transformer handle tokens when an expert is overloaded?

When an expert exceeds its capacity, the implementation randomly shuffles the token indices assigned to that expert using `torch.randperm` and truncates the list to the capacity limit (lines 119‑130). Dropped tokens are stored separately and later passed through the layer unchanged, meaning they bypass expert computation but preserve their original values and gradients. This prevents any single expert from becoming a computational bottleneck while maintaining the full sequence length in the final output.

### What is the difference between scaling probabilities and using a straight-through estimator?

When `is_scale_prob=True`, the expert outputs are multiplied directly by the routing probability `route_prob_max`, scaling the contribution of each token by its selection confidence (line 144). When `is_scale_prob=False`, the code uses a straight-through estimator: `route_prob_max / route_prob_max.detach()` (line 148). This division yields a value of 1.0 in the forward pass (preserving the magnitude) but allows gradients to flow backward through the routing decision by treating the denominator as a constant during backpropagation.

### Where is the Switch Transformer routing mechanism implemented in the labml repository?

The core routing logic resides in [`labml_nn/transformers/switch/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/switch/__init__.py), specifically within the `SwitchFeedForward` class. This file contains the router construction (lines 79‑80), probability computation (line 96), expert selection (line 100), capacity management (lines 108‑130), and output scaling (lines 144‑148). The training experiment and load-balancing loss implementation are located in [`labml_nn/transformers/switch/experiment.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/transformers/switch/experiment.py), which demonstrates how to integrate the routing statistics into a complete training pipeline.