How Capsule Networks Implement Dynamic Routing in PyTorch

Capsule Networks implement dynamic routing through an iterative routing-by-agreement algorithm where lower-level capsules project learned votes to higher-level capsules, and coupling coefficients update dynamically based on the scalar product agreement between votes and outputs, as realized in the Router class within labml_nn/capsule_networks/__init__.py.

The dynamic routing mechanism introduced in Dynamic Routing Between Capsules (Sabour et al., 2017) replaces scalar-based pooling operations with vector-based iterative consensus. In the labmlai/annotated_deep_learning_paper_implementations repository, this theoretical framework materializes as a clean PyTorch module that transforms primary capsule inputs into class-specific digit capsules through learned transformations and agreement-based coefficient updates.

The Mathematical Foundation of Routing-by-Agreement

Capsule Networks dynamic routing establishes part-whole relationships through an iterative process that measures the agreement between predicted poses and actual outputs. The algorithm computes coupling coefficients that determine how much each lower-level capsule contributes to higher-level capsules, updating these coefficients based on the similarity between predictions and final outputs.

Vote Projection and Transformation

For each lower-level capsule (i) and higher-level capsule (j), the network computes a prediction vector (\hat{u}{j|i}) by multiplying the lower-level capsule's pose vector (u_i) with a learned weight matrix (W{ij}). This projection generates the "vote" that lower-level capsule (i) casts for the pose of higher-level capsule (j).

Iterative Agreement Computation

The routing process iteratively refines routing logits (b_{ij}) (log prior probabilities) that determine coupling coefficients (c_{ij}) via softmax normalization. For each iteration, the algorithm:

  1. Computes weighted sums (s_j) of incoming votes
  2. Applies a squash non-linearity to produce output capsules (v_j)
  3. Calculates agreement (a_{ij}) as the scalar product (v_j \cdot \hat{u}_{j|i})
  4. Updates routing logits: (b_{ij} \leftarrow b_{ij} + a_{ij})

Implementation Breakdown of the Router Class

The Router class in labml_nn/capsule_networks/__init__.py encodes the complete dynamic routing algorithm. The implementation follows the mathematical specification precisely, using PyTorch operations to handle batched tensor computations efficiently.

Vote Projection via Einstein Summation (lines 96-99)

The first step transforms lower-level capsule inputs into prediction votes for higher-level capsules. The implementation uses torch.einsum to perform the batch matrix multiplication between learned weights and input capsules:

u_hat = torch.einsum('ijnm,bin->bijm', self.weight, u)

This operation produces a tensor of shape (batch, in_caps, out_caps, out_d), where each element represents the vote from a specific input capsule to a specific output capsule. The Einstein summation efficiently handles the four-dimensional weight tensor W_{ij} mapping from input dimension to output dimension.

Iterative Routing with Dynamic Coupling Coefficients (lines 14-29)

The routing logits (b_{ij}) initialize to zero for every capsule pair (lines 14-15). The algorithm then executes a fixed number of iterations (defaulting to three in the MNIST configuration) to refine the routing decisions:

Coupling Coefficient Computation (lines 20-21):

c = F.softmax(b, dim=2)

The softmax normalizes routing logits across output capsules, producing coupling coefficients (c_{ij}) that sum to one for each input capsule.

Weighted Aggregation and Non-Linearity (lines 22-25):

s = torch.einsum('bij,bijm->bjm', c, u_hat)
v = squash(s)

The weighted sum (s_j) aggregates all votes scaled by their coupling coefficients. The squash function then applies the non-linear scaling that preserves vector direction while enforcing length constraints on capsule activations.

Agreement Update (lines 26-29):

a = torch.einsum('bjm,bijm->bij', v, u_hat)
b = b + a

The scalar product (a_{ij}) measures the agreement between the output capsule (v_j) and the prediction vote (\hat{u}_{j|i}). High agreement increases the routing logit, strengthening the coupling coefficient in subsequent iterations.

After completing all iterations, the router returns the final squashed vectors (v_j) as the higher-level capsule activations (line 31).

Integrating Dynamic Routing into a Full Model

The MNIST experiment in labml_nn/capsule_networks/mnist.py demonstrates practical integration of the Router class within a complete architecture. The model transforms convolutional features into primary capsules, then applies dynamic routing to produce digit capsules.

Primary to Digit Capsule Transformation

The architecture processes input images through convolutional layers to generate 32 feature maps of 6×6 spatial dimensions, with each location containing an 8-dimensional capsule. This produces 1152 primary capsules (32 × 6 × 6) that feed into the routing layer:


# Primary capsule generation

x = self.conv2(F.relu(self.conv1(data)))                 # → [B, 32*8, 6, 6]

caps = x.view(x.shape[0], 8, 32*6*6).permute(0, 2, 1)   # → [B, 1152, 8]

caps = self.squash(caps)                                # squash primary capsules

# Dynamic routing to digit capsules

digit_caps = self.digit_capsules(caps)                  # → [B, 10, 16]

The digit_capsules module (instantiated as a Router with three iterations) transforms the 1152 primary capsules into 10 digit capsules, each with 16 dimensions.

Obtaining Class Predictions from Capsule Lengths

Following the paper's specification, the length of each output capsule vector encodes the probability that the corresponding class is present in the input:


# Length of each digit capsule encodes the probability of the class

probs = (digit_caps ** 2).sum(dim=-1).sqrt()   # shape: [B, 10]

predicted_class = probs.argmax(dim=-1)         # class index with highest probability

Practical Code Examples

Minimal Routing Layer Setup

To instantiate a routing layer independent of the full MNIST model:

import torch
from labml_nn.capsule_networks import Router, Squash

# 100 lower-level capsules, each 8-dimensional

in_caps = 100
in_dim  = 8

# 10 higher-level capsules, each 16-dimensional

out_caps = 10
out_dim  = 16

# three routing iterations (as in the paper)

router = Router(in_caps, out_caps, in_dim, out_dim, iterations=3)

# fake input: batch of 32 samples

x = torch.randn(32, in_caps, in_dim)

# forward pass – returns the output capsules (shape: [32, 10, 16])

digit_caps = router(x)

Key Source Files

  • labml_nn/capsule_networks/__init__.py: Contains the Router class implementing dynamic routing, the Squash activation function, and MarginLoss for training.
  • labml_nn/capsule_networks/mnist.py: End-to-end MNIST experiment that wires the routing layer into a full model, defining the training loop and evaluation metrics.
  • labml_nn/capsule_networks/mnist.ipynb: Interactive Colab notebook for training and visualizing the Capsule Network on MNIST.

Summary

  • Dynamic routing in Capsule Networks implements an iterative consensus mechanism where coupling coefficients update based on agreement between predicted and actual capsule poses.
  • The Router class in labml_nn/capsule_networks/__init__.py encodes this algorithm using torch.einsum for efficient vote projection and batched matrix operations for the iterative refinement process.
  • Routing logits initialize to zero and update via scalar product agreement between output capsules and prediction votes across multiple iterations (typically three).
  • The implementation transforms 1152 primary capsules (8-dimensional) into 10 digit capsules (16-dimensional) in the MNIST example, with the final capsule lengths serving as class probabilities.

Frequently Asked Questions

What is the purpose of dynamic routing in Capsule Networks?

Dynamic routing replaces max-pooling operations in traditional CNNs with a mechanism that preserves spatial hierarchies and pose information. By routing lower-level capsule votes to higher-level capsules based on agreement, the network ensures that spatial relationships and transformations are explicitly encoded in the vector outputs, making the architecture equivariant to transformations rather than simply invariant.

How many routing iterations are typically used?

The repository defaults to three routing iterations for the MNIST experiment, matching the original paper's configuration. Increasing iterations improves routing precision but increases computational cost linearly. The implementation allows configuring this via the iterations parameter in the Router class constructor.

What is the difference between routing logits and coupling coefficients?

Routing logits ((b_{ij})) are the log prior probabilities that initialize to zero and accumulate agreement scores across iterations. Coupling coefficients ((c_{ij})) are the softmax-normalized versions of these logits, representing the actual probability distribution determining how much each lower-level capsule contributes to each higher-level capsule. The logits store the "memory" of agreement, while the coefficients represent the current routing decision.

Why does the implementation use Einstein summation for vote projection?

The torch.einsum operation efficiently handles the four-dimensional tensor contraction required to map input capsules to output capsules across batch dimensions. The equation 'ijnm,bin->bijm' explicitly encodes the transformation from input dimension n to output dimension m for every pair of input capsule i and output capsule j, eliminating the need for explicit reshaping and batch matrix multiplication loops while maintaining mathematical clarity.

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 →