Vision Transformer (ViT)
Be able to explain patch embeddings and train a small ViT.
Prerequisites
- DTransformers — the architecturerequired
- EImage preprocessing and augmentationrequired
Intuition
A transformer wants a sequence of tokens. An image is a grid of pixels. ViT's solution is almost provocatively simple:
- Cut the image into square tiles — 16 × 16 pixels is standard.
- Flatten each tile into a vector and run it through a linear layer. Now every tile is a «token».
- Add a learnt positional embedding, since the transformer otherwise does not know where the tiles were.
- Add an extra
[CLS]token to gather up the whole. - Run an entirely ordinary transformer encoder.
- Classify from the
[CLS]token.
A 224 × 224 image with 16 × 16 patches becomes 14 × 14 = 196 tokens. That is a short sequence by language model standards.
The surprising result: it works better than a CNN — but only with enough data.
Formal
Inductive bias is the central difference.
A CNN has two assumptions built into the architecture: locality (neighbouring pixels belong together) and translation equivariance (the same filter everywhere). Those assumptions are true of images, and they are free — the network does not have to learn them.
A ViT has almost no assumptions. Attention connects any patch to any other already in the first layer. That means it has to learn that neighbouring patches belong together — which takes data.
| Training data | The winner |
|---|---|
| ~1 000 images | the CNN, by a wide margin |
| ImageNet-1k (1.3 M) | even, the CNN slightly ahead without tricks |
| ImageNet-21k (14 M) | the ViT |
| JFT-300M | the ViT, clearly |
The conclusion is not «ViTs are better» but that inductive bias can be replaced by data — and that it is a bad deal when data is scarce. With a small dataset: use a pretrained ViT and fine-tune it, or take a CNN.
Tricks that make a ViT manageable on small datasets:
| Trick | Effect |
|---|---|
| Heavy augmentation (RandAugment, Mixup, CutMix) | the largest single effect |
| Distillation from a CNN (DeiT) | brings in the CNN's bias via the teacher |
| Smaller patches on small images (4 × 4 on CIFAR) | more tokens, more local information |
| Learning rate warm-up + heavy weight decay | ViTs are more sensitive than CNNs |
The computational cost is quadratic in the number of patches. Doubling the resolution gives four times as many patches and sixteen times more expensive attention — hence hierarchical variants such as Swin, which do attention locally within windows.
Code
import torch, torch.nn as nn
class PatchEmbed(nn.Module):
"""Patching and projection are the same thing as a convolution with stride = the patch size."""
def __init__(self, image=32, patch=4, channels=3, d=192):
super().__init__()
self.n = (image // patch) ** 2
self.proj = nn.Conv2d(channels, d, kernel_size=patch, stride=patch)
def forward(self, x):
return self.proj(x).flatten(2).transpose(1, 2) # (B, n, d)
class ViT(nn.Module):
def __init__(self, n_classes=10, image=32, patch=4, d=192, depth=6, heads=3):
super().__init__()
self.patch = PatchEmbed(image, patch, 3, d)
self.cls = nn.Parameter(torch.zeros(1, 1, d))
self.pos = nn.Parameter(torch.randn(1, self.patch.n + 1, d) * 0.02)
layer = nn.TransformerEncoderLayer(d, heads, d * 4, dropout=0.1,
batch_first=True, norm_first=True,
activation="gelu")
self.enc = nn.TransformerEncoder(layer, depth)
self.norm = nn.LayerNorm(d)
self.head = nn.Linear(d, n_classes)
def forward(self, x):
z = self.patch(x)
z = torch.cat([self.cls.expand(z.size(0), -1, -1), z], dim=1) + self.pos
return self.head(self.norm(self.enc(z))[:, 0])
m = ViT()
print(m(torch.randn(2, 3, 32, 32)).shape) # torch.Size([2, 10])
print(sum(p.numel() for p in m.parameters()) / 1e6, "M parameters") # ~2.7 M
print("tokens:", m.patch.n + 1) # 65
Two details worth noticing:
nn.Conv2dwithstride=patchIS the patch embedding. Cutting out tiles and projecting them linearly is exactly the same operation as a convolution without overlap. It is more efficient and less code.norm_first=True(pre-LN). ViTs train considerably more stably with normalisation before the attention rather than after — the same experience as in language models.
Without heavy augmentation this model will overfit CIFAR-10 within a few epochs. That is not a bug in the code but the point of the section on inductive bias.
Mastery means
- Explains how an image becomes a sequence of tokens
- Compares ViT and CNN in inductive bias and data requirements
- Trains a small ViT on a small dataset
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — An Image is Worth 16x16 Words (ViT) — arXiv (open access; licence per article)
- arXiv — Training data-efficient image transformers (DeiT) — arXiv (open access; licence per article)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0