Skip to content
AI-grafen
EUniversityTransformer architecture· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

Vision Transformer (ViT)

Be able to explain patch embeddings and train a small ViT.

Prerequisites

Intuition

A transformer wants a sequence of tokens. An image is a grid of pixels. ViT's solution is almost provocatively simple:

  1. Cut the image into square tiles — 16 × 16 pixels is standard.
  2. Flatten each tile into a vector and run it through a linear layer. Now every tile is a «token».
  3. Add a learnt positional embedding, since the transformer otherwise does not know where the tiles were.
  4. Add an extra [CLS] token to gather up the whole.
  5. Run an entirely ordinary transformer encoder.
  6. 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 dataThe winner
~1 000 imagesthe CNN, by a wide margin
ImageNet-1k (1.3 M)even, the CNN slightly ahead without tricks
ImageNet-21k (14 M)the ViT
JFT-300Mthe 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:

TrickEffect
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 decayViTs 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.Conv2d with stride=patch IS 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

All the sources and licences