Vision & Multimodal Transformers
Images are sliced into a grid of non-overlapping patches, projected into linear embeddings, and processed through standard transformer blocks alongside text tokens.
Why Does This Exist?
For nearly a decade, Convolutional Neural Networks (CNNs) dominated computer vision through hard-coded inductive biases: translation equivariance (a feature detector produces the same response regardless of where an object appears) and two-dimensional spatial locality (pixels close together matter much more than pixels far apart). While these biases allow CNNs to learn sample-efficiently on modest datasets like ImageNet with 1.2 million images, they impose rigid computational constraints. A kernel cannot connect pixels on opposite edges of an image without stacking dozens of layers, diluting semantic information across repeated downsampling operations.
When industrial pretraining reached tens or hundreds of millions of images, these spatial inductive biases transformed from helpful guardrails into performance ceilings. Natural language processing had already converged on the standard Transformer architecture, which possesses almost no locality bias and scales compute and parameters with empirical power-law predictability.
Vision Transformers (ViT) eliminated custom vision architectures by discarding convolutional filters entirely. An image is cut into a fixed grid of square patches, each patch is flattened and linearly projected into a vector embedding, and the resulting sequence is handed directly to a vanilla Transformer encoder. By converting spatial pixels into token sequences, the exact same attention engine can simultaneously process image patches and text tokens, unlocking modern unified vision-language foundation models. Prerequisite understanding relies on backpropagation across residual layers.
Think of It Like This
A mosaic artist reading tiles as a sentence
Imagine a stained-glass window depicting a harbor scene. A convolutional inspector examines the window through a tiny magnifying glass, sliding it two centimeters at a time across the glass, noting local edges and color gradients, and then repeating the scan at lower resolution.
A Vision Transformer does something completely different. It cuts the window into a clean grid of square glass tiles. It numbers each tile from 1 to 196 to preserve its original position on the wall, and lays them out in a straight line on a long examination table.
Beside the tiles, the artist places a blank notebook labeled [CLS]. Because the artist possesses global vision across the entire table, tile 1 (the morning sun in the top-left corner) can immediately exchange light and color information with tile 190 (the ship's reflection in the bottom water) in a single glance. No tile is forced to communicate solely with its immediate physical neighbor. When reading text captions alongside the window, words and glass tiles sit on the exact same table, allowing the word "mast" to attend directly to tile 84 containing the ship's rigging.
How It Actually Works
Patch Extraction and Linear Projection
Given an input RGB image where is image height, is width, and is color channels, the model divides the spatial canvas into non-overlapping square patches of size . The total number of patches, which forms the effective sequence length , is:
Each patch is flattened into a one-dimensional vector of dimension . A learnable linear projection matrix maps each flattened patch into the transformer's latent model dimension :
where:
- is a prepended learnable classification token whose state at the output layer serves as the aggregate image representation.
- is a learnable 1D position embedding matrix added element-wise to retain spatial ordering, since standard self-attention is permutation-invariant.
The combined token sequence passes through identical Transformer encoder layers. Each layer consists of Multi-Head Self-Attention () and a Multi-Layer Perceptron () with GeLU activations, wrapped in pre-LayerNorm () operations and residual connections:
In multimodal architectures such as CLIP, ALIGN, or Flamingo, visual patch tokens and textual subword tokens are projected into a shared semantic space. Cross-attention layers allow text tokens to attend directly to visual patch representations:
where visual queries retrieve semantic attributes from textual keys and values .
Worked Example
Consider a standard ViT-Base processing a single image with resolution , , , patch size , and embedding dimension :
-
Patch count and input vector sizing: Each individual patch contains raw pixel values.
-
Patch projection: Projection matrix has shape . Multiplying 196 flattened vectors of length 768 by yields a patch embedding tensor of shape .
-
Prepend class token and position embeddings: Adding the single learnable vector produces sequence length . The resulting tensor has shape . The 1D position embedding matrix of shape is added element-wise.
-
Self-attention matrix computation: With 12 attention heads, each head operates on key-query dimension . For each head, the attention map has dimensions . Every forward pass computes attention coefficients per head, allowing every patch to observe all 195 other patches across the canvas.
Code
import torchimport torch.nn as nn
class PatchEmbedding(nn.Module): """Slices an image into non-overlapping patches and projects to embedding dim.""" def __init__(self, img_size: int = 224, patch_size: int = 16, in_channels: int = 3, embed_dim: int = 768): super().__init__() self.num_patches = (img_size // patch_size) ** 2 # A 2D convolution with kernel_size=patch_size and stride=patch_size performs # patch extraction and linear projection in a single optimized GPU operation self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x: torch.Tensor) -> torch.Tensor: # x: (batch_size, 3, 224, 224) x = self.proj(x) # -> (batch_size, 768, 14, 14) x = x.flatten(2) # -> (batch_size, 768, 196) return x.transpose(1, 2) # -> (batch_size, 196, 768)
class MinimalViT(nn.Module): def __init__(self, img_size: int = 224, patch_size: int = 16, embed_dim: int = 768, num_classes: int = 1000): super().__init__() self.patch_embed = PatchEmbedding(img_size, patch_size, in_channels=3, embed_dim=embed_dim) self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, 1 + self.patch_embed.num_patches, embed_dim)) encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=12, dim_feedforward=3072, activation="gelu", batch_first=True ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=1) self.head = nn.Linear(embed_dim, num_classes)
def forward(self, x: torch.Tensor) -> torch.Tensor: b = x.shape[0] tokens = self.patch_embed(x) # (b, 196, 768) cls_tokens = self.cls_token.expand(b, -1, -1) # (b, 1, 768) x = torch.cat((cls_tokens, tokens), dim=1) # (b, 197, 768) x = x + self.pos_embed # (b, 197, 768) encoded = self.encoder(x) # (b, 197, 768) logits = self.head(encoded[:, 0]) # Take [CLS] output -> (b, 1000) return logits
# Verificationmodel = MinimalViT()dummy_image = torch.randn(2, 3, 224, 224)output = model(dummy_image)print("Output shape:", output.shape)# -> Output shape: torch.Size([2, 1000])Watch Out For
Interpolating position embeddings across input resolutions
Practitioners commonly pretrain a Vision Transformer on images and then fine-tune on higher resolutions like or to boost fine-grained classification accuracy. However, increasing image dimensions from to increases patch count from to .
Directly loading pretrained weights throws a runtime shape mismatch error because was instantiated with length , while the model now requires positions. The standard fix is bicubic interpolation: separate the embedding from the spatial grid embeddings, reshape the remaining vectors into a grid, interpolate spatially to a grid, flatten back to vectors, and re-attach the token. Skipping spatial 2D grid reshaping and performing naive 1D linear interpolation destroys the two-dimensional spatial geometric relationships learned during pretraining.
The Quick Version
- Vision Transformers bypass convolutional kernels by slicing an image into a regular 2D grid of non-overlapping patches (), projecting each into an embedding vector.
- A learnable token and 1D position embeddings are added to the patch sequence before passing through standard Multi-Head Self-Attention layers.
- Self-attention enables immediate global context across any pair of patches regardless of spatial distance, scaling effectively to multimodal architectures where vision and text share a unified sequence representation.