Computer Vision

Vision transformers

A vision transformer cuts a picture into squares and lets every square look directly at every other square, instead of sliding a small window across the image.

On this page 7
  1. What came before, and what it could not do
  2. How it works
  3. Why it needs so much data
  4. Where you have already seen it
  5. What is honestly true about them
  6. Remember this
  7. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A vision transformer cuts a picture into small squares and lets every square look at every other square at once.

Think of doing a jigsaw puzzle on a table. You do not work strictly left to right. You pick up one piece and scan the whole table for anything that matches it. The same shade of sky. The same edge of a roof.

Your eyes jump straight across the table. The piece you need might be a hand's width away, and you find it without touching everything in between.

That jump, straight from one piece to a far-away piece, is the thing this whole lesson is about.

What came before, and what it could not do

A convolutional network slides a small window across the picture, over and over. Each window sees a tiny neighbourhood. Stack many layers and each one sees a little more.

That works beautifully, and it has one built-in limit. Connecting the top-left corner of a photograph with the bottom-right corner takes many layers. Information travels a small step at a time.

Sometimes that is exactly what you want. A dog's ear really is near a dog's head, and looking locally is a sensible assumption.

Sometimes it is not. Telling a beach photo from a swimming pool photo can need a direct comparison. The sky at the top, against the ground at the bottom. A transformer does that in one step.

How it works

Four moves, in order.

Cut the picture into equal squares. A standard setup uses squares of sixteen dots across and sixteen down. A normal-sized input then gives one hundred and ninety-six squares. Each square is called a patch.

Turn each square into a list of numbers. The grid of colours becomes a flat list. That list is called a token, borrowing the word from language models, where a token is one chunk of text.

Add a note saying where it came from. Cutting a picture into squares and shuffling them loses all sense of position. So each token gets a position embedding, a small set of numbers saying where it came from. It records "I was the fourth square in the second row".

Let every square look at every other square. Each token compares itself with all the others and decides which ones matter to it. That comparison step is attention, the same idea used to read a sentence and decide which words matter most.

   photo
     |
     +--> cut into 196 squares
              |
              +--> each square becomes a list of numbers  (a token)
                        |
                        +--> add "where I came from"      (position)
                                  |
                                  +--> every token looks at every token   (attention)
                                            |
                                            +--> one extra token collects
                                                 the summary -> "military uniform"

That last box is a real trick. One extra token is added that contains no picture at all. Its whole job is to gather from the others and carry the final answer. It is called the class token.

Why it needs so much data

Here is the honest heart of it, and it is worth reading twice.

A convolutional network is told, by its own design, that nearby dots belong together. It is also told that a pattern found in one corner is the same pattern in another. Nobody had to teach it that. It is baked into the shape of the model.

A transformer is told none of that. All it gets is a bag of squares and permission to compare any square with any other. It has to discover that neighbouring squares are related, from examples alone.

That freedom is the strength and the cost. Given a truly enormous pile of pictures, it learns better rules than the ones we would have written. Given a small pile, it flounders where a convolutional network would have done fine.

This is why vision transformers arrived late. The first version needed hundreds of millions of training images to beat the older approach.

Where you have already seen it

  • Searching your photo gallery by typing a description. The text-and-image models behind that are usually transformers.
  • Image generators. Type a sentence, get a picture. Transformers sit in most of that pipeline.
  • "Find this object" tools that outline anything you click on.
  • Modern medical and satellite systems, where the newest models are transformer-based.

What is honestly true about them

They are not automatically better. In 2022 a team rebuilt a convolutional network using the modern training tricks that transformers had brought with them. It matched the transformer at the same compute budget. Much of what looked like an architecture win turned out to be a training-recipe win.

They are expensive at high resolution. Every square compares itself with every other square. Double the number of squares and the comparison work goes up four times, not two. This is the practical reason plain vision transformers are not used directly on very large images.

The small examples in this lesson run on a laptop. Training one does not. Reproducing a published vision transformer takes many graphics cards for many days. That is not something a laptop can do, and no amount of patience changes it. What you can do on a laptop is use one somebody else trained, which is what almost everyone does.

Remember this

  • A vision transformer cuts the picture into patches and treats each as a token.
  • Every token can look at every other token directly, in one step.
  • It has no built-in idea that nearby dots belong together, so it needs far more data to learn that.
  • Using a pretrained one is easy. Training one from nothing is not a laptop activity.

What to learn next

  • Transformers — the architecture this borrows wholesale, explained from the language side.
  • Attention — the single operation everything above is built on.
  • Image segmentation — where transformer backbones now dominate.

Developer — Code and libraries.

Four short programs. The first three use nothing but NumPy and build the whole idea from scratch. The fourth inspects a real vision transformer without downloading a single byte of weights.

Setup

bash
pip install numpy torch torchvision matplotlib

Only the first program needs NumPy alone. Download sizes, stated up front:

WhatSizeWhen
PyTorch, CPU builda few hundred MBonce, at install
ViT-B/16 architecturenothingbuilt locally from code
ViT-B/16 pretrained weights330 MBonly for the optional last block

Step one: cut the picture into patches

patchify.py
import numpy as np

IMG, PATCH = 8, 2
GRID = IMG // PATCH                    # 4 patches across, 4 down

# A tiny image: mostly dark on the left, bright on the right, plus one odd patch.
img = np.full((IMG, IMG), 30.0)
img[:, 4:] = 200.0
img[2:4, 0:2] = 200.0                  # one bright patch stranded in the dark half

print("the image:")
for row in img:
    print("".join("#" if v > 100 else "." for v in row))

# Cut it into patches, then flatten each patch into one vector.
patches = (img.reshape(GRID, PATCH, GRID, PATCH)
              .transpose(0, 2, 1, 3)
              .reshape(GRID * GRID, PATCH * PATCH))
print()
print("patches:", patches.shape, "->", GRID * GRID, "patches of", PATCH * PATCH, "numbers each")
print("patch 0 (top-left) :", patches[0])
print("patch 3 (top-right):", patches[3])
Output
the image:
....####
....####
##..####
##..####
....####
....####
....####
....####

patches: (16, 4) -> 16 patches of 4 numbers each
patch 0 (top-left) : [30. 30. 30. 30.]
patch 3 (top-right): [200. 200. 200. 200.]

The reshape, transpose, reshape dance is the whole of patch extraction. Read it slowly. Split each axis into "which patch" and "which pixel inside the patch". Move the two patch axes to the front, then flatten.

Notice what has happened conceptually. A two-dimensional picture became a flat list of sixteen vectors. Every spatial relationship — which patch is next to which — has been thrown away. Hold that thought; step three is about getting it back.

Step two: let every patch look at every other patch

attention.py
import numpy as np

IMG, PATCH = 8, 2
GRID = IMG // PATCH

img = np.full((IMG, IMG), 30.0)
img[:, 4:] = 200.0
img[2:4, 0:2] = 200.0

patches = (img.reshape(GRID, PATCH, GRID, PATCH)
              .transpose(0, 2, 1, 3)
              .reshape(GRID * GRID, PATCH * PATCH))

# Real ViTs apply layer normalisation before attention. Centre and scale, as it does.
tokens = (patches - patches.mean()) / patches.std()

def softmax(x, axis=-1):
    e = np.exp(x - x.max(axis=axis, keepdims=True))
    return e / e.sum(axis=axis, keepdims=True)

d = tokens.shape[1]
attention = softmax(tokens @ tokens.T / np.sqrt(d))

print("attention matrix:", attention.shape)
print("every row sums to:", round(float(attention.sum(axis=1)[0]), 6))
print()
for query in (0, 3, 4):
    kind = "bright" if patches[query][0] > 100 else "dark  "
    print(f"how much patch {query} ({kind}) looks at each of the 16 patches:")
    for row in attention[query].reshape(GRID, GRID):
        print("   " + "  ".join(f"{v:.3f}" for v in row))
    print()
Output
attention matrix: (16, 16)
every row sums to: 1.0

how much patch 0 (dark  ) looks at each of the 16 patches:
   0.141  0.141  0.001  0.001
   0.001  0.141  0.001  0.001
   0.141  0.141  0.001  0.001
   0.141  0.141  0.001  0.001

how much patch 3 (bright) looks at each of the 16 patches:
   0.003  0.003  0.109  0.109
   0.109  0.003  0.109  0.109
   0.003  0.003  0.109  0.109
   0.003  0.003  0.109  0.109

how much patch 4 (bright) looks at each of the 16 patches:
   0.003  0.003  0.109  0.109
   0.109  0.003  0.109  0.109
   0.003  0.003  0.109  0.109
   0.003  0.003  0.109  0.109

Stop and read the second grid, because this is the point of the entire lesson.

Patch 3 sits in the top-right corner. In its attention grid, the value at row one, column zero is 0.109 — high. That is patch 4, the stranded bright patch on the middle left, on the far side of the image.

They found each other in one operation. There is no stack of layers, no gradual spreading of information. Every patch compared itself with every patch, and matching content lit up regardless of distance. A convolutional network needs many layers to connect those two positions.

Also notice each attention row sums to exactly 1. Attention distributes a fixed budget across all patches, so paying more attention somewhere means paying less elsewhere. That property is what makes the numbers readable as "how much it looked at each one".

Step three: without position, it is a bag of squares

Look again at the grids for patch 3 and patch 4. They are identical. Two patches at opposite ends of the picture, treated as interchangeable, because their contents match and nothing else was recorded.

position.py
import numpy as np

IMG, PATCH = 8, 2
GRID = IMG // PATCH

img = np.full((IMG, IMG), 30.0)
img[:, 4:] = 200.0
img[2:4, 0:2] = 200.0
patches = (img.reshape(GRID, PATCH, GRID, PATCH)
              .transpose(0, 2, 1, 3).reshape(GRID * GRID, PATCH * PATCH))
tokens = (patches - patches.mean()) / patches.std()

def softmax(x):
    e = np.exp(x - x.max(axis=-1, keepdims=True))
    return e / e.sum(axis=-1, keepdims=True)

def attend(t):
    return softmax(t @ t.T / np.sqrt(t.shape[1]))

print("WITHOUT position information")
a = attend(tokens)
print("  patch 3 and patch 4 have identical content.")
print("  are their attention rows identical?", np.allclose(a[3], a[4]))
print()

# One number per patch saying where it sits. Real ViTs learn this vector.
rows = np.repeat(np.arange(GRID), GRID)
cols = np.tile(np.arange(GRID), GRID)
position = np.stack([rows, cols, rows * 0.0, cols * 0.0], axis=1) * 0.5

print("WITH position added to every patch")
b = attend(tokens + position)
print("  are their attention rows identical?", np.allclose(b[3], b[4]))
print()
print("  patch 3 (top-right corner) now looks at:")
for row in b[3].reshape(GRID, GRID):
    print("   " + "  ".join(f"{v:.3f}" for v in row))
print()
print("  patch 4 (middle-left) now looks at:")
for row in b[4].reshape(GRID, GRID):
    print("   " + "  ".join(f"{v:.3f}" for v in row))
Output
WITHOUT position information
  patch 3 and patch 4 have identical content.
  are their attention rows identical? True

WITH position added to every patch
  are their attention rows identical? False

  patch 3 (top-right corner) now looks at:
   0.000  0.000  0.060  0.110
   0.023  0.000  0.075  0.137
   0.000  0.000  0.094  0.170
   0.000  0.000  0.117  0.212

  patch 4 (middle-left) now looks at:
   0.001  0.001  0.058  0.072
   0.053  0.001  0.082  0.102
   0.001  0.002  0.115  0.144
   0.002  0.002  0.163  0.203

Attention on its own has no idea where anything is. Shuffle the patches and it produces the same answers in a different order. That is a mathematical property, not a bug.

Adding a position vector to each token is what fixes it, and it is the entire mechanism. In a real vision transformer that vector is learned rather than written by hand. It is added to the token before the first block.

That also explains a practical annoyance you will hit. The position vectors are learned for one specific number of patches. Feed the model a different image size and the count changes, so the stored positions no longer line up. Every library works around this by stretching the learned position grid to the new shape. The model loses a little accuracy in the process.

Step four: a real vision transformer, no download

weights=None builds the architecture from code and fetches nothing.

inspect_vit.py
import torch
from torchvision.models import vit_b_16

# weights=None builds the architecture without downloading anything.
model = vit_b_16(weights=None).eval()

print("patch embedding layer:", model.conv_proj)
print("class token :", tuple(model.class_token.shape))
print("position    :", tuple(model.encoder.pos_embedding.shape))
print("encoder blocks:", len(model.encoder.layers))
print("parameters  :", f"{sum(p.numel() for p in model.parameters()):,}")
print()

x = torch.zeros(1, 3, 224, 224)
tokens = model._process_input(x)          # cut into patches and embed them
print("after patch embedding:", tuple(tokens.shape))
tokens = torch.cat([model.class_token.expand(1, -1, -1), tokens], dim=1)
print("after adding the class token:", tuple(tokens.shape))
print()
print("patches across:", 224 // 16, " patches down:", 224 // 16,
      " total:", (224 // 16) ** 2)
print("attention scores per head, per layer:", (224 // 16) ** 2 + 1, "x", (224 // 16) ** 2 + 1)
Output
patch embedding layer: Conv2d(3, 768, kernel_size=(16, 16), stride=(16, 16))
class token : (1, 1, 768)
position    : (1, 197, 768)
encoder blocks: 12
parameters  : 86,567,656

after patch embedding: (1, 196, 768)
after adding the class token: (1, 197, 768)

patches across: 14  patches down: 14  total: 196
attention scores per head, per layer: 197 x 197

Four things in that output are worth a minute each.

The patch embedding is a convolution. Conv2d(3, 768, kernel_size=16, stride=16). Kernel size equals stride, so the windows never overlap — it lands on each patch exactly once and never revisits a pixel. "Cut into patches and multiply each by a weight matrix" and "one non-overlapping convolution" are the same operation. Every implementation does it this way because it is faster.

196 patches, then 197 tokens. The extra one is the class token. It is a learned vector with no image content, prepended to the sequence, and the classifier reads only its output. Everything the model concluded has to be funnelled through that one token.

Position embedding shape is (1, 197, 768). One vector per token position, learned during training, including a position for the class token.

Every layer builds a 197 by 197 table of attention scores, per head. That table is what grows with the square of the patch count. At 224 pixels it is trivial. At 1024 pixels it would be 4097 by 4097, per head, per layer. That is why plain vision transformers are not applied directly to large images.

Optional: run the pretrained model

This downloads 330 MB. Everything above already works without it. Skip this on mobile data.

classify_vit.py
import torch
import matplotlib.cbook as cbook
from torchvision.io import decode_image
from torchvision.models import vit_b_16, ViT_B_16_Weights

image = decode_image(cbook.get_sample_data("grace_hopper.jpg").name)
weights = ViT_B_16_Weights.IMAGENET1K_V1
print("weights file:", weights.meta["_file_size"], "MB")
print("ImageNet top-1:", weights.meta["_metrics"]["ImageNet-1K"]["acc@1"])

model = vit_b_16(weights=weights).eval()
with torch.no_grad():
    logits = model(weights.transforms()(image).unsqueeze(0))[0]

probs = logits.softmax(0)
top = probs.topk(5)
for score, index in zip(top.values, top.indices):
    print(f"  {weights.meta['categories'][index]:28s} {float(score):.3f}")
Output
weights file: 330.285 MB
ImageNet top-1: 81.072
  military uniform             0.684
  bearskin                     0.072
  bow tie                      0.064
  suit                         0.036
  comic book                   0.015

"Military uniform" at 0.684, on a portrait of a naval officer. The runners-up are a tall fur military hat and a bow tie, both plausible reads of the same region. This is a well-behaved model failing gracefully rather than confidently.

Now the comparison nobody puts on the slide

These figures come straight out of torchvision's own metadata, all measured on ImageNet-1k.

ModelWeightsParametersTop-1
ResNet-50 (V2 recipe)97.8 MB25.6M80.86
ConvNeXt-Tiny109.1 MB28.6M82.52
ViT-B/16330.3 MB86.6M81.07
ViT-L/161161.0 MB304.3M79.66

Read those four rows carefully, because they contradict the usual story.

ViT-B/16 uses roughly three and a half times the parameters of ResNet-50 for about the same accuracy. On this benchmark, at this scale, the transformer is not winning.

ConvNeXt-Tiny beats both, with fewer parameters than the transformer and a modernised convolutional design. That is the result from the what-is-computer-vision lesson, shown here in downloadable numbers.

ViT-L/16 is worse than ViT-B/16, despite three and a half times the parameters and a gigabyte of weights. This is the data-hunger problem, measured. These checkpoints were trained on ImageNet-1k alone, and the larger transformer has more capacity than that dataset can supervise. Pre-trained on a far larger corpus, the ordering flips.

That single row is the most useful thing on this page. A transformer's advantage is a function of how much data it saw, not of its architecture.

Common mistakes

Training a ViT from scratch on a small dataset. It will underperform a ResNet and it will look like your code is broken. Start from pretrained weights, or use a convolutional model, or use a recipe explicitly designed for small data such as DeiT's.

Changing the input resolution and expecting it to work. The position embeddings were learned for one patch count. Libraries interpolate them, and accuracy drops. If you need a different resolution, fine-tune at it.

Comparing parameter counts as if they were compute. A ViT's attention cost depends on patch count, not only on parameters. Two models with equal parameters can differ several-fold in latency at high resolution.

Reading attention maps as explanations. An attention weight says which token was mixed in, not which token caused the decision. Attention rollout and gradient-based methods exist; treating a raw attention map as an explanation is a well-documented mistake.

Forgetting weights.transforms(). ViT checkpoints have their own resize, crop and normalisation. Get any of them wrong and accuracy falls with no error raised.

Try it yourself

Change PATCH from 2 to 4 in attention.py. The image is still eight by eight, so you now get four patches instead of sixteen. Predict two things before running it. How many rows does the attention matrix have? And is the stranded bright region still visible as its own patch?

It is not. The stranded bright region now falls inside a patch that is mostly dark, and it stops being a token of its own. Its pixels are still in there — the flattened patch vector holds every one of them — but attention works at patch granularity. Run it. Patch 0 spends most of its attention on itself and on the other dark patch, and almost ignores the two bright ones.

That is patch size in one experiment. Nothing smaller than a patch can be attended to on its own. Smaller patches give the model finer things to reason about, and cost far more. The attention table grows with the square of the patch count.

Then run inspect_vit.py with vit_b_32 instead of vit_b_16. Forty-nine tokens instead of one hundred and ninety-six, so the attention table shrinks about sixteen-fold. Now look up its published accuracy. ViT_B_32_Weights.IMAGENET1K_V1 reports 75.9 top-1 against ViT-B/16's 81.1, with slightly more parameters, because the patch projection grew. Coarser patches are cheaper in attention and dearer in everything that matters.

What to learn next

  • Transformers — the architecture this borrows wholesale, explained from the language side.
  • Attention — the single operation everything above is built on.
  • Image segmentation — where transformer backbones now dominate.

Researcher — Mathematics and papers.

The architecture

Dosovitskiy et al. (2021) applied the standard Transformer encoder to images with essentially no vision-specific modification.

Patch embedding. Reshape x in R^(H x W x C) into N flattened patches and project:

text
N = H * W / P^2
z_0 = [ x_class ; x_p^1 E ; x_p^2 E ; ... ; x_p^N E ] + E_pos
  • P — patch side length, 16 in ViT-B/16.
  • x_p^i in R^(P^2 * C) — patch i, flattened.
  • E in R^(P^2 C x D) — the linear projection; D is the model width, 768 for ViT-B.
  • x_class in R^D — a learned token prepended to the sequence.
  • E_pos in R^((N+1) x D) — learned position embeddings.

For 224 x 224 input with P = 16, N = 196 and the sequence length is 197. The projection E is implemented as a Conv2d with kernel_size = stride = P, which is numerically identical and far faster.

Encoder block, pre-norm ordering:

text
z'_l = MSA( LN(z_{l-1}) ) + z_{l-1}
z_l  = MLP( LN(z'_l) )    + z'_l
  • LN — layer normalisation.
  • MSA — multi-head self-attention.
  • The MLP is two linear layers with GELU, expansion ratio 4.

Attention:

text
Attention(Q, K, V) = softmax( Q K^T / sqrt(d_k) ) V
  • Q = z W_Q, K = z W_K, V = z W_V, each in R^(N x d_k).
  • d_k = D / h — per-head dimension; h is the number of heads, 12 in ViT-B.
  • The sqrt(d_k) divisor keeps the dot products from saturating the softmax as d_k grows.

Configurations:

ModelLayersWidth DHeadsParams
ViT-B/16127681286M
ViT-L/1624102416307M
ViT-H/1432128016632M

Cost

Per layer, with sequence length N and width D:

text
attention:  O(N^2 * D)     the  N x N  score matrix, per head
MLP:        O(N * D^2)     the two projections

The crossover is at N ≈ D. For ViT-B/16 at 224 pixels, N = 197 and D = 768, so the MLP dominates and attention is not the bottleneck. Double the resolution to 448 and N = 784, now exceeding D, and attention takes over — cost in the attention term rises sixteen-fold for a four-fold increase in pixels.

This is why plain ViT does not scale to dense prediction at high resolution, and why every hierarchical variant restricts attention somehow.

Memory for the score matrix alone is O(h * N^2) per layer in the naive implementation. FlashAttention (Dao et al., 2022) computes exact attention without materialising it, tiling the computation to stay in on-chip SRAM. It is an IO-complexity result, not an approximation, and it is the reason long-sequence attention became affordable.

Position embeddings

The original paper compared learned 1-D, learned 2-D and relative embeddings, and found little difference — 1-D learned is the default. The result is less surprising than it first appears. A learned embedding can encode 2-D structure without being told the layout, and the paper's own visualisations show exactly that recovery.

Practical consequences:

  • Resolution change requires interpolation. E_pos is learned for a specific N. Changing the input size means resampling the position grid, usually bicubically, with a measurable accuracy cost. Fine-tuning after interpolation recovers most of it.
  • Relative position bias (used in Swin) generalises across resolutions more gracefully.
  • Rotary embeddings, standard in language models, are increasingly used in vision and remove the fixed-length constraint.

Data scale is the whole story

The paper's central experiment: ViT trained on ImageNet-1k underperforms a comparable BiT ResNet; on ImageNet-21k they are close; on JFT-300M the ViT wins outright and scales better.

The interpretation is about inductive bias. Convolution supplies locality and translation equivariance as hard architectural constraints. A transformer must learn both from data. When data is plentiful, learned structure beats imposed structure. When it is scarce, the imposed structure is worth roughly an order of magnitude of data.

Three lines of work reduced the data requirement substantially:

DeiT (Touvron et al., 2021) reached competitive ImageNet-1k accuracy with no external data, using heavy augmentation, stochastic depth, repeated augmentation, and a distillation token supervised by a convolutional teacher. The distillation token is the interesting part — the ViT learns from a CNN's inductive bias without inheriting its architecture.

AugReg (Steiner et al., 2021) systematically measured augmentation and regularisation against dataset size, and showed that a well-regularised ViT on ImageNet-21k matches a ViT trained on much more data. Their headline conclusion is uncomfortable and useful: many published architecture comparisons are confounded by training recipe.

Self-supervised pre-training. MAE (He et al., 2022) masks 75 percent of patches and reconstructs them; the asymmetric encoder sees only visible patches, making pre-training cheap. DINO and DINOv2 (Caron et al., 2021; Oquab et al., 2023) use self-distillation and produce features whose attention maps segment objects without any segmentation labels. DINOv2 features support linear probes that approach full fine-tuning on many dense tasks.

Hierarchical variants

Plain ViT has one resolution throughout, which suits classification and not dense prediction.

Swin Transformer (Liu et al., 2021) computes attention within local windows and shifts the window partition between consecutive blocks so information crosses boundaries. Cost becomes linear in image size rather than quadratic, and the multi-scale feature maps drop into existing detection and segmentation necks.

PVT, MViT and others reduce sequence length progressively by pooling keys and values.

ViTDet (Li et al., 2022) is the counterargument: plain, non-hierarchical ViT backbones with a simple feature pyramid built afterwards are competitive for detection, given strong pre-training such as MAE. Hierarchy may be a convenience rather than a necessity.

The ConvNeXt correction

Liu et al. (2022) took a ResNet-50 and applied, one at a time, the design and training choices that came with transformers: the training recipe, depthwise convolutions with large kernels, an inverted bottleneck, fewer normalisation layers, GELU, LayerNorm. The resulting pure ConvNet matched Swin at equal FLOPs across classification, detection and segmentation.

The honest reading is not "convolutions won". It is that the 2020-to-2021 transformer results conflated architecture with training methodology, and once the methodology was transferred, the architectural gap largely closed. Any claim that attention is intrinsically better for vision has to engage with this paper.

A known artefact

Darcet et al. (2024), Vision Transformers Need Registers, identified high-norm tokens appearing in low-information background patches of trained ViTs, including DINOv2. The model appears to repurpose uninformative patches as scratch storage for global information, which corrupts the feature map for dense tasks and makes attention visualisations misleading.

The fix is small and effective: append a few extra learnable "register" tokens that are discarded at output. Attention maps become clean and dense-task performance improves. It is a useful reminder that trained networks develop internal conventions nobody designed, and that reading attention maps as explanation is hazardous.

Papers

Where this stands

The plain ViT is now infrastructure rather than a result. It is the default backbone for multimodal models, for open-vocabulary detection and segmentation, and for self-supervised feature learning, mainly because it shares an architecture and a scaling story with language models. Whether it is the best classifier at a given FLOP budget is no longer the interesting question.

Three problems remain genuinely open.

High-resolution dense prediction. Quadratic attention still constrains what is affordable. Window attention, token merging and linear-attention approximations all trade something away, and none is a clean win.

Small-data regimes. Most real projects have thousands of images, not millions. Self-supervised pre-training plus a linear probe is currently the best answer. It also presumes a pre-trained backbone matching your domain, which for medical, satellite and industrial imagery frequently does not exist.

Interpretability. Attention maps are not explanations, the register-token result shows that trained ViTs hide state in unexpected places, and faithful attribution for transformers remains unsolved. Anyone deploying one where a wrong answer has a cost should assume they cannot explain its decisions, and design the surrounding system accordingly.

What to learn next

  • Transformers — the architecture this borrows wholesale, explained from the language side.
  • Attention — the single operation everything above is built on.
  • Image segmentation — where transformer backbones now dominate.