Computer Vision2015intermediate11 min read

Spatial Transformer Networks

شبكات المحوِّل المكاني

Jaderberg, M. · Simonyan, K. · Zisserman, A. · Kavukcuoglu, K. — NeurIPS

The problem

CNNs rely on small max-pooling windows (typically 2×2) for . This means a digit rotated 45° or shifted to the corner looks like a completely different input to the network. Building to large transformations requires stacking many pooling layers, wasting depth and capacity on a problem that has nothing to do with recognition itself.

The contribution

The : a differentiable module with three parts — a that predicts transformation parameters from the input, a grid generator that maps output coordinates back to input coordinates, and a differentiable sampler that reads pixel values using . The entire module is trained end-to-end with , requires no extra supervision, and can be inserted anywhere inside a .

The impact

Spatial Transformers introduced the idea that neural networks can learn to spatially manipulate their own maps. This concept directly influenced deformable convolutions, mechanisms in detection (DETR), and in 3D vision. It showed that geometric reasoning can be learned end-to-end rather than hard-coded.

Imagine you're scanning a stack of handwritten envelopes. Some addresses are upside down, some are tilted, and some are crammed into a corner. A traditional CNN is like a clerk who can only read straight text — every crooked envelope is a puzzle.

The Spatial Transformer gives the clerk a robot hand that grabs each envelope, rotates it upright, centers the address, and hands it back — all before reading. The clerk never learns to read upside-down text; instead, the hand learns to straighten everything so reading is easy.

The key trick: the hand doesn't follow fixed rules. It looks at the envelope, decides what adjustment to make, and does it differently for every single one.

The problem: pooling gives only local invariance

CNNs use max-pooling layers to tolerate small spatial shifts. A 2×2 pooling window means a feature can shift by one pixel without affecting the output. But large transformations — 30° rotations, 2× scale changes, or significant translations — cannot be absorbed by stacking a few pooling layers.

The traditional fix is : train on rotated, scaled, and translated copies of each image. This helps, but it doesn't solve the problem — it just asks the network to memorize every possible pose rather than reason about geometry. The wastes capacity storing redundant representations of the same object at different poses.

Open in Lab
Left: max-pooling absorbs a 1-pixel shift but fails at rotation. Right: a spatial transformer rotates the entire input to an upright canonical pose.
The demo wakes as you arrive…

The architecture: three parts working together

The spatial transformer has three components that execute in sequence, like a three-step assembly line:

1. Localization network — looks at the input and predicts what transformation to apply. It can be any small CNN or fully-connected network. Its final layer outputs θ, the transformation parameters (6 numbers for an ). Think of it as the "eyes" that assess how the image is oriented.

2. Grid generator — takes θ and builds a : for every pixel in the output, it computes where to look in the input. This is the "plan" — a map from target coordinates to source coordinates. For an affine transformation, this is a simple matrix multiplication.

3. Sampler — executes the plan. For each output pixel, it reads the input at the computed source location using bilinear interpolation. Because bilinear interpolation is differentiable, gradients flow backward through the sampler to the localization network, so the whole system learns end-to-end.

Open in Lab
Click "Transform" to watch the three stages in sequence: the localization network reads the input, the grid generator builds the sampling map, and the sampler produces the corrected output.
The demo wakes as you arrive…

The affine transformation: 6 numbers that reshape an image

The heart of the grid generator is a coordinate mapping. For every pixel location (xit,yit)(x^t_i, y^t_i) in the output (target), we compute where to read from in the input (source). In the affine case, this is a 2×3 matrix multiplication.

Before seeing the formula, think of it like this: you have a sheet of graph paper (the output). For each square on your sheet, the affine matrix tells you: "go look at this spot on the original image." Scaling shrinks or stretches the grid. Rotation spins it. Translation slides it. All are encoded in just 6 numbers.

(xisyis)=(θ11θ12θ13θ21θ22θ23)(xityit1)\begin{pmatrix} x^s_i \\ y^s_i \end{pmatrix} = \begin{pmatrix} \theta_{11} & \theta_{12} & \theta_{13} \\ \theta_{21} & \theta_{22} & \theta_{23} \end{pmatrix} \begin{pmatrix} x^t_i \\ y^t_i \\ 1 \end{pmatrix}
Affine grid mapping — from output to input coordinates — (xit,yit)(x^t_i, y^t_i) = target pixel in the output · (xis,yis)(x^s_i, y^s_i) = where to sample from the input · θ₁₃, θ₂₃ = translation · θ₁₁, θ₂₂ = scale · θ₁₂, θ₂₁ = rotation/shear · The localization network predicts all 6 θ values.
Open in Lab
Drag the sliders to change the 6 affine parameters and watch how the sampling grid deforms. The output always shows what the network "sees" after transformation.
The demo wakes as you arrive…

Differentiable sampling: the key insight

The grid generator gives us source coordinates (xis,yis)(x^s_i, y^s_i), but these are typically between pixel locations — not at exact integer positions. How do you read a pixel value at position (3.7, 5.2)?

Bilinear interpolation: you look at the four nearest integer pixels and blend them, weighted by how close each one is. Position (3.7, 5.2) would take 70% of the value from column 4 and 30% from column 3, and similarly blend rows 5 and 6.

The crucial property: this blending is differentiable. You can compute ∂V/∂xis\partial V / \partial x^s_i — how the output changes as you nudge the sample location — and chain that back through the grid generator to the localization network. This is what allows the network to learn where to look, using nothing but the final classification loss.

Vic=∑nH∑mWUnmcmax⁡(0,1−∣xis−m∣) max⁡(0,1−∣yis−n∣)V^c_i = \sum_n^H \sum_m^W U^c_{nm} \max(0, 1 - |x^s_i - m|) \, \max(0, 1 - |y^s_i - n|)
Bilinear sampling — differentiable pixel reading — For each output pixel VicV^c_i, sum contributions from all input pixels UnmcU^c_{nm}, weighted by proximity to the sample point. In practice, only the 4 nearest pixels contribute (the max clips everything else to zero). This formula is differentiable with respect to both UU and (xis,yis)(x^s_i, y^s_i).
Open in Lab
Drag the sample point (red dot) across the input grid. Watch how the four neighboring pixels contribute to the output value, with weights shown as bar heights.
The demo wakes as you arrive…

How gradients flow: end-to-end learning

The gradient chain is the magic that makes spatial transformers self-supervised. No one tells the network how to transform the input — the classification loss itself teaches the localization network what transformation helps recognition.

The chain works backward: the classification loss produces a gradient on the output feature map V. The sampler formula is differentiable with respect to both the input feature map U and the grid coordinates (xs,ys)(x^s, y^s). The grid coordinates are a differentiable function of θ (a matrix multiply). And θ is the output of the localization network — a standard neural network that backpropagation handles naturally.

∂L∂θ=∂L∂V⋅∂V∂(xs,ys)⋅∂(xs,ys)∂θ\frac{\partial \mathcal{L}}{\partial \theta} = \frac{\partial \mathcal{L}}{\partial V} \cdot \frac{\partial V}{\partial (x^s, y^s)} \cdot \frac{\partial (x^s, y^s)}{\partial \theta}
The gradient chain through the spatial transformer — Loss → output feature map → sample coordinates → transformation parameters. Each link is differentiable, so standard backpropagation trains the entire module. No reinforcement learning, no separate supervision — just the task loss.
Open in Lab
Click "Backpropagate" to animate the gradient flowing from the classification loss back through the sampler, grid generator, and localization network.
The demo wakes as you arrive…

Beyond affine: projective and thin plate spline

The affine transformation is just the simplest option. The spatial transformer framework supports any differentiable transformation:

  • Affine (6 parameters): rotation, scale, translation, shear. Lines stay parallel.

  • Projective (8 parameters): adds perspective. Parallel lines can converge, like a photo taken from an angle.

  • (TPS) (variable parameters): non-rigid warping controlled by a set of control points. Can straighten curved handwriting or correct local distortions. The most powerful but most parameter-heavy option.

The choice depends on what transformations the data actually contains. For rotated digits, affine is sufficient. For elastic deformations in handwriting, TPS excels.

Open in Lab
Compare affine, projective, and thin plate spline transformations side by side. Toggle each to see how the grid deforms and what the output looks like.
The demo wakes as you arrive…

Placement: input, mid-network, or parallel

A spatial transformer can be inserted at multiple positions in a CNN, and each placement serves a different purpose:

At the input — acts like an intelligent pre-processor. The localization network sees raw pixels and learns to crop, rotate, and scale the whole image. This is the simplest and most common setup.

Between convolutional layers — transforms feature maps rather than raw images. Deeper feature maps encode richer semantics, so the localization network can make more informed decisions. The paper's best result on street view house numbers uses four spatial transformers stacked at increasing depths.

Multiple in parallel — each transformer focuses on a different region or object part. For bird classification, one transformer learns to crop the head and another the body — without any part annotation supervision.

The same idea in code

A minimal spatial transformer in NumPypython

Simplified to show the idea — not the real implementation.

import numpy as np

def affine_grid(theta, H, W):
    """Build a sampling grid from a 2×3 affine matrix.
    Returns source coordinates for each output pixel."""
    # Normalized coordinates: -1 to 1
    yt = np.linspace(-1, 1, H)
    xt = np.linspace(-1, 1, W)
    xt, yt = np.meshgrid(xt, yt)
    ones = np.ones_like(xt)
    # Stack into (3, H*W) homogeneous coordinates
    grid = np.stack([xt.ravel(), yt.ravel(), ones.ravel()])
    # Apply affine: (2, 3) @ (3, H*W) -> (2, H*W)
    source = theta @ grid   # source coords in input space
    xs = source[0].reshape(H, W)
    ys = source[1].reshape(H, W)
    return xs, ys

def bilinear_sample(U, xs, ys):
    """Sample from image U at fractional coordinates (xs, ys).
    This is the differentiable part — gradients flow through here."""
    H, W = U.shape[:2]
    # Convert from [-1,1] to pixel indices
    xs = (xs + 1) * 0.5 * (W - 1)
    ys = (ys + 1) * 0.5 * (H - 1)
    x0, y0 = np.floor(xs).astype(int), np.floor(ys).astype(int)
    x1, y1 = x0 + 1, y0 + 1
    # Clip to image bounds
    x0, x1 = np.clip(x0, 0, W-1), np.clip(x1, 0, W-1)
    y0, y1 = np.clip(y0, 0, H-1), np.clip(y1, 0, H-1)
    # Bilinear weights
    wa = (x1 - xs) * (y1 - ys)
    wb = (xs - x0) * (y1 - ys)
    wc = (x1 - xs) * (ys - y0)
    wd = (xs - x0) * (ys - y0)
    return wa[...,None]*U[y0,x0] + wb[...,None]*U[y0,x1] \
         + wc[...,None]*U[y1,x0] + wd[...,None]*U[y1,x1]

# Usage: localization_net predicts theta, then:
# xs, ys = affine_grid(theta, H_out, W_out)
# output = bilinear_sample(input_image, xs, ys)

Results: what spatial transformers learn

The experiments demonstrate three key findings:

  • Distorted MNIST: Adding a spatial transformer to a simple fully-connected network (ST-FCN) matches the of a CNN with max-pooling, purely by learning to normalize poses. Adding it to a CNN (ST-CNN) reduces error further — from 0.8% to 0.5% on rotated/translated/scaled digits.

  • Street View House Numbers (SVHN): Using four stacked spatial transformers inside a CNN achieves 3.6% sequence error, beating the previous state of the art (3.9%) which used model ensembles and Monte Carlo averaging. The spatial transformers learn to crop and zoom into individual digit regions.

  • Fine-grained bird classification (CUB-200): Two parallel spatial transformers spontaneously learn to act as a head detector and a body detector — without any part annotations. This improves accuracy from 82.3% to 84.1%.

Why it changed everything

Before spatial transformers, geometric reasoning in neural networks was hard-coded through pooling and data augmentation. The STN paper showed that a network can learn to manipulate its own feature maps spatially — and that this ability is differentiable, composable, and requires no extra labels.

This opened a family of ideas: if you can learn where to look (attention), you can learn how to deform (deformable convolutions), what region to attend to (detection transformers like DETR), and even how to render 3D scenes (differentiable rendering).

  1. 2015

    Spatial Transformer Networks

    Jaderberg et al. introduce the STN — a differentiable module that learns input-dependent spatial transformations end-to-end. No extra supervision needed.

  2. 2017

    Deformable Convolutional Networks

    Dai et al. generalize the idea: instead of transforming the whole feature map, each convolution kernel learns per-pixel offsets. Irregular receptive fields adapt to object shape.

  3. 2020

    DETR — Detection Transformer

    Carion et al. use attention as a learned spatial query mechanism for object detection, eliminating hand-designed components like anchors and NMS. A direct descendant of the "learn where to look" philosophy.

  4. 2020

    NeRF — Neural Radiance Fields

    Mildenhall et al. push differentiable spatial reasoning to 3D: a neural network learns a continuous 3D scene, and differentiable rendering produces novel views. The core idea — differentiable coordinate mapping — traces directly to STN.

The spatial transformer's legacy is not the module itself — few production systems use an explicit STN today. Its legacy is the principle: geometric reasoning in neural networks should be learned, not engineered. Every , every attention-based detector, every differentiable renderer inherits this idea.

CitationJaderberg, Simonyan, Zisserman, Kavukcuoglu. Spatial Transformer Networks. NeurIPS, 2015.

Terms in this paper