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

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


# 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:

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

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:

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

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

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

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

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

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.

Practical Implementation Example

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

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:

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 (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, 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, which demonstrates how to integrate the routing statistics into a complete training pipeline.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →