Modern AI Engineering

Lesson 16.2 · 25 min

Vision Transformers: Applying Self-Attention to Image Patches

Transformers were built to read sentences — so how can the very same architecture look at a photo and say 'that is a cracked blender jar'?

In short: A Vision Transformer (ViT) treats an image like a sentence: it cuts the image into small square patches, turns each patch into a vector (a 'visual word'), adds a special [CLS] token and position information, and runs the sequence through a standard Transformer encoder. The final [CLS] vector summarises the whole image and a small head turns it into class scores. ViTs need more data than CNNs to learn well, but they scale very well and are now the default image encoder inside most multimodal models.

The big picture

In earlier lessons we saw the Transformer: a stack of self-attention and feed-forward layers that reads a sequence of token vectors and lets every token look at every other token. It took over language processing. In 2020, a Google team asked a simple question: what if we feed an image to an almost unchanged Transformer? Their paper, 'An Image is Worth 16x16 Words', introduced the Vision Transformer (ViT).

The trick is all in the input. A Transformer needs a sequence of vectors. An image is a grid of pixels. So ViT turns the grid into a sequence: chop the picture into small squares called patches, flatten each one, and project it to a vector. From there on, the model is the same encoder we already know.

Think of it like a jigsaw puzzle Imagine cutting a photo into 196 puzzle pieces and laying them in a row. Each piece alone shows only a tiny bit — a patch of white, an edge of glass. To recognise the picture, we must look at the pieces together and remember where each came from. Self-attention is how the pieces 'look at' each other, and position embeddings are the numbers on the back telling us where each piece belongs.

We will decode ViT in six steps, using a concrete example: the standard ViT-Base/16 model classifying a 224×224 photo from our appliance shop. 'Base' is the model size; '/16' means 16×16-pixel patches.

Decoding step 1: splitting the image into patches

Why not just treat each pixel as a token? Because self-attention compares every token with every other token, so its cost grows with the square of the sequence length. A 224×224 image has 50,176 pixels; attention over all of them would mean about 2.5 billion pairs per head per layer. That is far too expensive.

Patches fix this. With patch size P = 16, the image becomes (224/16) × (224/16) = 14 × 14 = 196 patches. Each patch is a small 16×16×3 block that we flatten (lay out as one long list) into 16 × 16 × 3 = 768 numbers. The patches do not overlap, and we read them in row order, left to right, top to bottom, like words on a page.

Pause and think: A ViT uses 14×14 patches on a 224×224 image. How many patches, and how many numbers per flattened RGB patch?

224/14 = 16, so 16 × 16 = 256 patches. Each patch holds 14 × 14 × 3 = 588 numbers.

Decoding step 2: patch embedding

A flattened patch is just raw pixel values. We need it to be a vector of the model's width D (the hidden size, 768 in ViT-Base) so the Transformer can work with it. ViT multiplies every flattened patch by the same learned matrix E of size 768 × D, plus a bias. The result is called a patch embedding — the visual equivalent of a word embedding.

Sharing the matrix matters: a vertical edge looks the same whether it appears in the top-left or bottom-right corner, so the same weights should detect it everywhere. During training, the rows of E learn to respond to useful local patterns such as edges, colour blobs and textures.

In real code it is a convolution Implementations usually write the 'split + flatten + linear' steps as a single convolution with kernel size 16 and stride 16. Because the kernel and stride are equal, each output position sees exactly one non-overlapping patch, which is mathematically the same as the linear projection.

Decoding step 3: the [CLS] token

After step 2 we have 196 vectors, one per patch. But to classify the image we want one vector that describes the whole picture. ViT borrows a trick from BERT: it puts an extra, learnable vector at the start of the sequence called the [CLS] token (short for 'classification'). The sequence is now 197 tokens long.

The [CLS] token holds no pixels. It starts as the same learned vector for every image. As it passes through the encoder layers, self-attention lets it gather information from all 196 patches. By the last layer, its vector has become a summary of the image, and that is the only vector the classifier reads.

An alternative: average pooling Instead of a [CLS] token, we can average the 196 output patch vectors (global average pooling). The original paper reported both can work well with tuned settings, and many later vision models use pooling. Either way, the goal is one vector for the whole image.

Decoding step 4: position embeddings

Self-attention has a blind spot: it treats its input as a set. If we shuffled the 196 patches, the attention math would give the same outputs, just shuffled. But in an image, location matters: sky usually sits above ground, and a crack 'at the base' must be at the bottom. So ViT adds a position embedding to every token: a learned vector, one per position (0 for [CLS], 1 to 196 for the patches).

Interestingly, these are plain 1-D positions (just 'patch number 37'), not row and column. The original paper tried 2-D-aware versions and found no significant gain: after training, the learned position vectors of patches in the same row or column became similar on their own. The model discovered the 2-D grid.

Changing the image size The position table has a fixed number of rows. If we fine-tune at a higher resolution (say 384×384 → 576 patches), there are no learned vectors for the new positions. The standard fix is to arrange the old vectors on their 2-D grid and interpolate them to the new grid size.

Decoding step 5: the Transformer encoder

Now the 197 vectors enter a standard Transformer encoder: a stack of L identical layers (L = 12 in ViT-Base). Each layer has two sub-blocks, each wrapped in a residual connection and preceded by LayerNorm (the 'pre-norm' layout):

Inside one encoder layer

  1. LayerNorm: Normalise each token vector so its numbers have a stable scale. This keeps training of deep stacks stable.
  2. Multi-head self-attention: Every token makes a query, key and value. Each token's output is a weighted mix of all values, weighted by how well its query matches each key. There is no causal mask: a patch can look at patches before and after it. ViT-Base uses 12 heads, so it can track 12 kinds of relationships at once.
  3. Add (residual): Add the attention output back to the input. The token keeps its own information and adds context from other patches.
  4. LayerNorm + MLP: A two-layer feed-forward network with a GELU activation, applied to each token independently (768 → 3072 → 768 in ViT-Base). This is where much of the per-token processing happens.
  5. Add (residual) again: Add the MLP output back. The result has the same shape, 197 × 768, and becomes the input of the next layer.

Because every patch can attend to every other patch from the very first layer, ViT has a global receptive field everywhere. A CNN, by contrast, only sees a small neighbourhood per layer and grows its view slowly with depth. Studies of trained ViTs found that some heads in early layers still attend locally — the model learns local processing when it is useful — while others look far across the image.

Decoding step 6: the classification head

After the last layer (and a final LayerNorm), we take only the output at position 0 — the [CLS] vector, 768 numbers — and feed it to a classification head. In pre-training the paper used a small MLP with one hidden layer; for fine-tuning on a new task, it is usually a single linear layer of size 768 × K, where K is the number of classes. A softmax turns the K scores into probabilities.

For our shop, K might be 5: 'intact jar', 'cracked jar', 'broken blade', 'burnt motor', 'other'. We would take a ViT pre-trained on a large dataset, replace its head with a new 768 × 5 layer, and fine-tune on a few thousand labelled support photos.

Putting it all together (with code)

Here is the whole pipeline on a tiny 8×8 image with 4×4 patches, a width of 16 and a single attention head, written in plain numpy so every step is visible. The weights are random (untrained), so the attention and the class probabilities are close to uniform — what matters is the flow of shapes.

tiny_vit.py

import numpy as np
rng = np.random.default_rng(0)
H = W = 8; C = 3; P = 4; D = 16           # tiny image, 4x4 patches, width 16
img = rng.random((H, W, C))                # an 8x8 RGB image
# Step 1: cut into non-overlapping P x P patches and flatten each one
patches = img.reshape(H // P, P, W // P, P, C).transpose(0, 2, 1, 3, 4)
patches = patches.reshape(-1, P * P * C)   # (N, P*P*C)
N = patches.shape[0]
print("patches:", patches.shape)           # 4 patches of 48 numbers
# Step 2: one shared linear layer turns each patch into a D-dim token
W_patch = rng.normal(scale=0.1, size=(P * P * C, D))
tokens = patches @ W_patch                  # (N, D)
# Steps 3-4: prepend a learnable [CLS] token, add position embeddings
cls = rng.normal(scale=0.1, size=(1, D))
pos = rng.normal(scale=0.1, size=(N + 1, D))
x = np.vstack([cls, tokens]) + pos          # (N+1, D)
print("sequence into encoder:", x.shape)
# Step 5: one self-attention head (no mask: every patch sees every patch)
Wq, Wk, Wv = (rng.normal(scale=0.3, size=(D, D)) for _ in range(3))
q, k, v = x @ Wq, x @ Wk, x @ Wv
s = q @ k.T / np.sqrt(D)
a = np.exp(s - s.max(1, keepdims=True)); a /= a.sum(1, keepdims=True)
x = x + a @ v                                # residual connection
print("CLS attends to [CLS, p1..p4]:", np.round(a[0], 2))
# Step 6: classification head reads ONLY the CLS output
W_head = rng.normal(scale=0.1, size=(D, 3))  # 3 classes
logits = x[0] @ W_head
probs = np.exp(logits) / np.exp(logits).sum()
print("class probabilities:", np.round(probs, 3))
# Real ViT-Base sizes: 224x224 image, 16x16 patches
n = (224 // 16) ** 2
print("ViT-B/16 patches:", n, "-> sequence length", n + 1,
"| values per patch:", 16 * 16 * 3)

Output:

patches: (4, 48)
sequence into encoder: (5, 16)
CLS attends to [CLS, p1..p4]: [0.2  0.2  0.19 0.2  0.21]
class probabilities: [0.332 0.347 0.321]
ViT-B/16 patches: 196 -> sequence length 197 | values per patch: 768
Standard ViT sizes from the original paper
ModelLayersHidden size DMLP sizeHeadsParameters
ViT-Base12768307212≈ 86M
ViT-Large241024409616≈ 307M
ViT-Huge321280512016≈ 632M

ViT vs CNN

Before ViT, image models were convolutional neural networks (CNNs) such as ResNet. A CNN slides small filters (say 3×3) across the image, so each layer only combines nearby pixels. This bakes in two assumptions, called inductive biases: locality (nearby pixels are related) and translation equivariance (a cat shifted right should produce shifted features). These are good assumptions for images, so CNNs learn well from modest data.

ViT drops most of these assumptions. It must learn from data that nearby patches matter. The original paper found that trained on ImageNet alone (about 1.3M images), ViT did worse than comparable ResNets; pre-trained on much larger datasets (14M to 300M images), it matched or beat them while using less compute to pre-train. Later work such as DeiT showed that strong data augmentation, regularisation and distillation let ViTs train well on ImageNet alone.

Pause and think: We have only 2,000 labelled support photos. Should we train a ViT from scratch?

No. With so little data a ViT from scratch will likely underperform because it lacks the CNN's built-in assumptions. Instead, fine-tune a ViT (or CNN) that was pre-trained on a large dataset — that brings the learned visual knowledge with it.

Where ViTs are used today ViT-style encoders are the image backbone in CLIP and in most vision-language chat models, which feed the patch vectors (after a projection) into an LLM. They are also used in image classification, as the encoder in segmentation and detection systems, in medical imaging research, and as the 'eyes' of robotics models.

Worked example, step by step

A good way to check that we understand ViT is to count its parameters by hand and see whether we land near the ≈ 86M listed for ViT-Base. We only need the sizes already given: P = 16, C = 3, D = 768, MLP size 3,072, 12 layers, 197 tokens.

Counting ViT-Base

  1. Patch embedding: E maps 768 flattened pixel values to D = 768, plus a bias: 768 × 768 + 768 = 590,592 parameters. It is shared by all 196 patches.
  2. [CLS] and positions: [CLS] is one vector: 768. The position table has 197 rows: 197 × 768 = 151,296. The whole input stage is 590,592 + 768 + 151,296 = 742,656.
  3. Attention in one layer: Query, key, value and output projections are each 768 × 768 + 768 = 590,592. Four of them: 2,362,368. The 12 heads split these matrices; they do not add parameters.
  4. MLP in one layer: 768 → 3,072 costs 768 × 3,072 + 3,072 = 2,362,368. 3,072 → 768 costs 3,072 × 768 + 768 = 2,360,064. Together: 4,722,432.
  5. One whole layer: Add two LayerNorms (2 × 768 numbers each = 3,072): 2,362,368 + 4,722,432 + 3,072 = 7,087,872.
  6. Twelve layers plus the input stage: 12 × 7,087,872 = 85,054,464. Add the input stage and the final LayerNorm (1,536): about 85.8M before the classification head. That matches the ≈ 86M in the table.
Where ViT-Base's parameters live (computed above)
PartParametersShare
Patch embedding + [CLS] + positions742,656under 1%
12 × attention28,348,416about 33%
12 × MLP56,669,184about 66%
LayerNorms38,400tiny

Two things stand out. The part that is special to images, the input stage, is under 1% of the model; everything else is an ordinary Transformer encoder. And the MLP blocks hold about twice as many parameters as attention. Notice also what does not depend on image size: only the position table does. That is why fine-tuning at a higher resolution needs no new weights except interpolated positions, yet still costs far more compute, because the number of token pairs grows.

Practice: try it yourself

We will cut a tiny numbered image into patches so we can see exactly which pixels land in which patch. Then we test the 'blind spot' from step 4: we shuffle the patches and check whether a one-head model notices, with and without position embeddings.

practice_patch_shuffle.py

import numpy as np
rng = np.random.default_rng(3)
# A 4x4 one-channel "image" whose pixels are numbered 0..15, patch size 2.
img = np.arange(16).reshape(4, 4)
P = 2
patches = img.reshape(2, P, 2, P).transpose(0, 2, 1, 3).reshape(-1, P * P)
for i, p in enumerate(patches):
print(f"patch {i}: pixels {p}")
D = 8
E = rng.normal(scale=0.5, size=(P * P, D))       # shared patch embedding
pos = rng.normal(scale=0.5, size=(4, D))         # one vector per position
Wq, Wk, Wv = (rng.normal(scale=0.3, size=(D, D)) for _ in range(3))
def image_vector(patch_rows, use_pos):
"""Embed patches, run one attention head, average into one vector."""
x = (patch_rows / 15.0) @ E                  # scale pixels to 0..1
if use_pos:
x = x + pos                              # slot i always gets pos[i]
s = (x @ Wq) @ (x @ Wk).T / np.sqrt(D)
a = np.exp(s - s.max(1, keepdims=True)); a /= a.sum(1, keepdims=True)
return (x + a @ x @ Wv).mean(axis=0)         # mean-pool the tokens
order = [3, 1, 0, 2]                             # scramble the patch order
shuffled = patches[order]
for use_pos in (False, True):
a = image_vector(patches, use_pos)
b = image_vector(shuffled, use_pos)
label = "with positions   " if use_pos else "without positions"
print(f"{label}: largest change after shuffling = {np.abs(a - b).max():.3f}")
for size in (224, 384):
n = (size // 16) ** 2 + 1
print(f"{size}x{size} image -> {n} tokens -> {n * n:,} attention pairs")

Output:

patch 0: pixels [0 1 4 5]
patch 1: pixels [2 3 6 7]
patch 2: pixels [ 8  9 12 13]
patch 3: pixels [10 11 14 15]
without positions: largest change after shuffling = 0.000
with positions   : largest change after shuffling = 0.042
224x224 image -> 197 tokens -> 38,809 attention pairs
384x384 image -> 577 tokens -> 332,929 attention pairs

Now change it:

  • Change order to [0, 1, 2, 3]. Predict both printed changes before running.
  • Multiply pos by 4 (use scale=2.0). Predict whether the 'with positions' change gets larger or smaller, and whether the 'without positions' line moves at all.
  • In the last loop, use size // 8 instead of size // 16. Predict the token count for 224×224 first, then how many times the pair count grows.

Pause and think: The code pools by averaging all tokens. If we read a [CLS] token instead, and still used no position embeddings, would shuffling the patches change the [CLS] output?

No. [CLS] gathers information through attention, and attention treats the patches as a set: each patch contributes the same value with the same weight wherever it sits in the sequence. The readout method does not fix the blind spot; only position information does.

Pause and think: We switch ViT-Base from 16×16 to 8×8 patches on 224×224 images. Does the model get many more parameters? What does change?

Hardly. The patch embedding actually shrinks (8·8·3 = 192 inputs instead of 768), and the position table grows from 197 to 785 rows, which is small. The encoder layers are untouched. What explodes is compute and memory: 785 tokens means about 16× more attention pairs per layer.

Quick summary

  • Split the image into non-overlapping P×P patches: N = (H/P)·(W/P), each flattened to P²·C numbers.
  • Project each patch with one shared linear layer to a D-dimensional patch embedding.
  • Prepend a learnable [CLS] token.
  • Add learned position embeddings so the model knows where each patch came from.
  • Run a standard Transformer encoder (no causal mask) over the sequence.
  • Feed the final [CLS] vector to a classification head.

Common mistakes Forgetting that attention cost grows quadratically with the number of patches (small patches on big images get expensive fast); training a ViT from scratch on a small dataset; feeding images at a different resolution without interpolating the position embeddings; and assuming attention maps are a faithful explanation of the decision.

Key takeaways

  • ViT turns an image into a sequence of patch tokens and processes it with a standard Transformer encoder.
  • Sequence length is (H/P)·(W/P) + 1; smaller patches mean more detail but quadratically more attention cost.
  • A learnable [CLS] token gathers a whole-image summary; position embeddings restore where each patch was.
  • ViT has weaker built-in image assumptions than a CNN, so it needs more data, but it scales very well.
  • ViT-style encoders are the standard 'eyes' of modern multimodal models.

Key terms

  • Patch: A small non-overlapping square of the image (e.g. 16×16 pixels) that becomes one token.
  • Patch embedding: The vector produced by projecting a flattened patch with a shared learned matrix.
  • [CLS] token: An extra learnable token placed first in the sequence whose final output summarises the image.
  • Position embedding: A learned vector added to each token so the model knows which position it came from.
  • Inductive bias: An assumption built into a model's design, such as locality in CNNs.
  • Receptive field: The region of the input that can influence a given feature.

← 16.1 Multimodal AI: Perceiving Text, Images, and Audio Together · 16.3 Image Embeddings: Encoding Visual Content as Vectors →