Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Attention in Vision Transformers (ViT)
In 2020 a group at Google took the transformer — an architecture designed for word sequences — deleted essentially nothing, chopped images into 16×16 squares, and fed them in as if they were tokens. No convolutions. No pooling. No hand-designed notion that nearby pixels belong together.
Trained on ImageNet's 1.28 million images it lost to a ResNet by several percentage points. That should have been the end of it. Instead they trained the same model on JFT-300M — 300 million images — and it reached 88.55% ImageNet top-1, beating the best convolutional network available at 87.54%, while using about four times less compute to get there.
The interesting part is why it lost on the small dataset and won on the large one. That gap is not a quirk of tuning. It is the exact price of the inductive biases a CNN builds in and a transformer does not, and understanding it tells you which architecture to reach for on your own data.
Two different assumptions about images
A convolutional layer hard-codes three beliefs about the visual world:
| Built-in assumption | How the architecture enforces it | When it is wrong |
|---|---|---|
| Locality — a pixel relates most to its neighbours | A 3×3 kernel sees nine pixels and nothing else | Relating two objects at opposite corners of the frame |
| Translation equivariance — a cat is a cat wherever it appears | The same kernel weights slide across every position | Tasks where absolute position matters (document layout, medical scans with fixed anatomy) |
| Hierarchy — edges compose into textures, textures into parts, parts into objects | Stacked layers with pooling, receptive field growing with depth | Rarely wrong; this one is genuinely good |
These assumptions are enormously useful when data is scarce. They are constraints the model does not have to learn, so every training example goes towards learning something else. But they are still constraints.
Quantify the locality one. With 3×3 convolutions and no downsampling, the receptive field grows by 2 pixels per layer. To let one output unit see across a 224-pixel image you need
Real CNNs shortcut this with striding and pooling — a ResNet-50 reaches a full-image receptive field by around layer 40 in the theoretical sense — but the effective receptive field, weighted by actual gradient contribution, is far smaller and roughly Gaussian around the centre. Long-range relationships in a CNN are always indirect, mediated by many intermediate layers.
Self-attention has the opposite profile. In layer 1, every patch attends to every other patch. The receptive field is the whole image immediately. What it does not know is that patch 5 and patch 6 are adjacent — it has to learn that from data, which is precisely why it needs so much of it.
A CNN is given the structure of images and spends its capacity on content. A transformer is given nothing and must spend some of its capacity rediscovering that structure — which is a bad trade on a million images and a good one on three hundred million.
Patch embedding: turning a picture into a sequence
Why not pixels
The obvious approach is one token per pixel. Do the arithmetic and it dies immediately. A 224×224 image has 50,176 pixels, so the attention matrix would be
per head, per layer. In fp16 that is about 5 GB for a single head of a single layer of a single image. Not feasible now and not feasible soon.
The patch grid
Instead, cut the image into a grid of non-overlapping squares and treat each square as one token. For image size H×W and patch size P:
With H=W=224 and P=16: N=2242/162=50176/256=196 patches, arranged 14×14.
Each patch is 16×16×3=768 raw numbers. Flatten it and project it linearly to the model dimension D:
For ViT-Base, D=768, so E is 768×768 — 589,824 weights plus 768 biases, giving 590,592 parameters in the patch embedding.
The attention matrix is now 197×197=38,809 entries (196 patches plus one extra token we will meet shortly). That is five orders of magnitude smaller than the pixel version, and entirely ordinary.
The implementation trick
Cutting patches, flattening and projecting is exactly a convolution with kernel size P and stride P:
1import torch2import torch.nn as nn34class PatchEmbedding(nn.Module):5 def __init__(self, img_size=224, patch_size=16, in_ch=3, embed_dim=768):6 super().__init__()7 assert img_size % patch_size == 0, "image size must divide by patch size"8 self.n_patches = (img_size // patch_size) ** 29 # kernel == stride == patch_size: each output position sees exactly one patch,10 # with no overlap. Mathematically identical to flatten-then-Linear.11 self.proj = nn.Conv2d(in_ch, embed_dim,12 kernel_size=patch_size, stride=patch_size)1314 def forward(self, x): # (B, 3, 224, 224)15 x = self.proj(x) # (B, 768, 14, 14)16 x = x.flatten(2) # (B, 768, 196)17 return x.transpose(1, 2) # (B, 196, 768)Using Conv2d here is not a sneaky reintroduction of convolutional bias. Because stride equals kernel size, no kernel ever spans two patches, so there is no weight sharing across overlapping regions and no locality assumption. It is a fast, well-optimised way to do a per-patch linear projection.
The assert matters more than it looks. If img_size is 225, the convolution silently drops the last row and column of pixels, your patch count is wrong, and the positional embedding table no longer lines up.
The architecture
The classification token
The transformer outputs one vector per token, but classification needs a single vector for the whole image. ViT borrows BERT's solution: prepend a learned vector, the [CLS] token, which belongs to no patch.
1self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) # 768 parameters2...3cls = self.cls_token.expand(B, -1, -1) # (B, 1, 768)4x = torch.cat([cls, x], dim=1) # (B, 197, 768)Because attention is symmetric in its access — every token sees every token — this extra position accumulates information from all 196 patches through the stack, and its final state feeds the classifier head. It is a learned aggregation, which turns out to work better than a fixed one because the model decides what to pool.
The alternative is global average pooling over the 196 patch outputs. Later work found the two perform about the same once the learning rate is retuned; the [CLS] token survives mostly by convention. It does have one genuine advantage: its attention row is directly interpretable as "which patches did the classifier look at".
Positional embeddings
Nothing in attention encodes position — permute the patch order and the outputs permute identically. Without positional information a ViT is a bag of patches, and a scrambled jigsaw would classify the same as the assembled picture.
ViT uses learned 1-D position embeddings: a table of 197×768=151,296 parameters, one row per sequence position, added to the patch embeddings. Not 2-D, not sinusoidal. The paper tested 2-D-aware variants and found no meaningful improvement — the model learns the grid structure by itself, and if you visualise the learned embeddings you find that nearby-in-2D positions end up with similar vectors even though nothing told them to.
The consequence to remember: the table has exactly 197 rows. Change the input resolution and the patch count changes, so the table no longer fits. Fine-tuning at higher resolution requires interpolating the embeddings — reshape the 196 patch rows into a 14×14 grid, bicubically resize to the new grid, flatten back:
1import torch.nn.functional as F23def interpolate_pos_embed(pos_embed, old_grid, new_grid):4 """pos_embed: (1, 1 + old_grid**2, D). Keeps the CLS row, resizes the rest."""5 cls_tok, patch_tok = pos_embed[:, :1], pos_embed[:, 1:]6 D = patch_tok.shape[-1]7 patch_tok = patch_tok.reshape(1, old_grid, old_grid, D).permute(0, 3, 1, 2)8 patch_tok = F.interpolate(patch_tok, size=(new_grid, new_grid),9 mode='bicubic', align_corners=False)10 patch_tok = patch_tok.permute(0, 2, 3, 1).reshape(1, new_grid ** 2, D)11 return torch.cat([cls_tok, patch_tok], dim=1)Forget this step and loading a 224-trained checkpoint at 384 resolution either raises a shape error or, if someone wrote a permissive loader, silently misaligns every position.
The full model
1class ViTBlock(nn.Module):2 """Pre-norm transformer block: LN inside the residual branch."""3 def __init__(self, dim, heads, mlp_ratio=4.0, dropout=0.0):4 super().__init__()5 self.norm1 = nn.LayerNorm(dim)6 self.attn = nn.MultiheadAttention(dim, heads, dropout=dropout,7 batch_first=True)8 self.norm2 = nn.LayerNorm(dim)9 hidden = int(dim * mlp_ratio)10 self.mlp = nn.Sequential(11 nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(dropout),12 nn.Linear(hidden, dim), nn.Dropout(dropout))1314 def forward(self, x, return_attn=False):15 h = self.norm1(x)16 a, w = self.attn(h, h, h,17 need_weights=return_attn,18 average_attn_weights=False) # keep per-head weights19 x = x + a20 x = x + self.mlp(self.norm2(x))21 return (x, w) if return_attn else (x, None)2223class VisionTransformer(nn.Module):24 def __init__(self, img_size=224, patch_size=16, in_ch=3, num_classes=1000,25 dim=768, depth=12, heads=12, mlp_ratio=4.0, dropout=0.0):26 super().__init__()27 self.patch_embed = PatchEmbedding(img_size, patch_size, in_ch, dim)28 n = self.patch_embed.n_patches2930 self.cls_token = nn.Parameter(torch.zeros(1, 1, dim))31 self.pos_embed = nn.Parameter(torch.zeros(1, n + 1, dim))32 self.pos_drop = nn.Dropout(dropout)3334 self.blocks = nn.ModuleList(35 ViTBlock(dim, heads, mlp_ratio, dropout) for _ in range(depth))36 self.norm = nn.LayerNorm(dim) # required with pre-norm blocks37 self.head = nn.Linear(dim, num_classes)3839 nn.init.trunc_normal_(self.pos_embed, std=0.02)40 nn.init.trunc_normal_(self.cls_token, std=0.02)4142 def forward(self, x, return_attn=False):43 B = x.shape[0]44 x = self.patch_embed(x) # (B, 196, D)45 x = torch.cat([self.cls_token.expand(B, -1, -1), x], 1) # (B, 197, D)46 x = self.pos_drop(x + self.pos_embed)4748 maps = []49 for blk in self.blocks:50 x, w = blk(x, return_attn)51 if return_attn:52 maps.append(w)5354 x = self.norm(x)55 logits = self.head(x[:, 0]) # the CLS position only56 return (logits, maps) if return_attn else logitsThe self.norm after the last block is not optional with pre-norm blocks. The residual stream is never normalised on the way out, so without it the features arriving at the classifier grow with depth and the logits come out badly scaled.
Sizes
| Model | Layers | D | Heads | MLP | Parameters |
|---|---|---|---|---|---|
| ViT-Base | 12 | 768 | 12 | 3072 | 86 M |
| ViT-Large | 24 | 1024 | 16 | 4096 | 307 M |
| ViT-Huge | 32 | 1280 | 16 | 5120 | 632 M |
Verify ViT-Base from its parts. Per block: attention is four 768×768 matrices with biases, 4(7682+768)=2,362,368; the MLP is 768⋅3072+3072+3072⋅768+768=4,722,432; two LayerNorms add 4×768=3,072. That is 7,087,872 per block, and 12×7,087,872=85,054,464. Add the patch embedding (590,592), position table (151,296) and CLS token (768) and you land just under 86 M.
Note again where the parameters are: the MLP holds twice what attention does. Two thirds of a ViT is feed-forward.
The naming convention ViT-B/16 means Base with patch size 16. The patch size is not a minor detail — it controls sequence length quadratically.
Visualising attention, and the bug everyone hits
The [CLS] row of the attention matrix is the natural thing to plot: 196 weights, one per patch, summing to 1, reshaped into a 14×14 heatmap.
1import matplotlib.pyplot as plt23@torch.no_grad()4def cls_attention_map(model, image, layer=-1, head=None, grid=14):5 model.eval()6 _, maps = model(image.unsqueeze(0), return_attn=True)7 w = maps[layer][0] # (heads, 197, 197)8 w = w.mean(0) if head is None else w[head]9 cls_row = w[0, 1:] # drop the CLS-to-CLS entry10 return cls_row.reshape(grid, grid).cpu()The bug. PyTorch's nn.MultiheadAttention defaults to average_attn_weights=True, which returns the mean over heads. Averaging 12 heads that attend to different things produces a bland, roughly uniform map, and people conclude the model is not attending to anything meaningful. It is — you have just blurred twelve distinct patterns together. Pass average_attn_weights=False and inspect heads individually.
Two further traps in the same area:
| Trap | What you see | Fix |
|---|---|---|
F.scaled_dot_product_attention is used internally | No weights at all — the fused kernel never materialises the matrix | Register a forward hook on the explicit path, or set the library flag that disables fusion for inspection |
| Forgetting to drop index 0 before reshaping | 197 values will not reshape to 14×14, or you silently shift every patch by one | Slice [1:] before reshape(14, 14) |
| Plotting only the last layer | Diffuse maps with strong sinks on a few patches | Use attention rollout, below — the last layer alone does not describe the whole path |
That last one deserves the maths. A single layer's attention tells you where information moved in that layer, but the CLS token at layer 12 is a mixture of things that already moved in layers 1 through 11. Attention rollout composes them. Account for the residual connection by mixing each layer's attention with the identity, renormalise, and multiply:
1@torch.no_grad()2def attention_rollout(maps):3 """maps: list of (B, heads, T, T). Returns (T, T) for the first example."""4 T = maps[0].shape[-1]5 result = torch.eye(T)6 for a in maps:7 a = a[0].mean(0) # average heads within a layer8 a = 0.5 * a + 0.5 * torch.eye(T) # account for the residual path9 a = a / a.sum(dim=-1, keepdim=True) # renormalise rows to sum to 110 result = a @ result11 return resultRollout maps are markedly cleaner: on an image of a dog on grass, the last-layer CLS row is diffuse, while the rollout concentrates on the animal's head and body.
One honest caveat. High attention weight means information flowed along that edge. It does not prove the prediction depended on it — the value vector could be near zero, and later MLPs can discard whatever arrived. Attention maps generate hypotheses; gradient-based attribution tests them.
Training in practice
The data threshold
| Pre-training data | Best CNN (BiT / ResNet) | ViT | Winner |
|---|---|---|---|
| ImageNet-1k (1.3 M) | ~77% | ~72–77% for ViT-B | CNN |
| ImageNet-21k (14 M) | ~85% | ~84–85% | roughly level |
| JFT-300M (300 M) | 87.54% (BiT-L) | 88.55% (ViT-H/14) | ViT |
The crossover sits somewhere around 10–100 million images. Below it, the CNN's built-in locality is worth more than the transformer's flexibility; above it, the transformer learns better structure than the one we designed by hand.
Practically, almost nobody trains a ViT from scratch. You start from a checkpoint pre-trained on ImageNet-21k or larger and fine-tune. That inverts the table above — a fine-tuned ViT-B/16 will typically beat a fine-tuned ResNet-50 on a 5,000-image custom dataset, because the pre-training already paid the data cost.
If you must train on a small dataset, DeiT showed it is possible: strong augmentation (RandAugment, Mixup, CutMix) and stochastic depth take ViT-B to 81.8% on ImageNet-1k alone, and distillation from a CNN teacher lifts that to about 83.4%. The distillation is the telling part — the CNN teacher is transferring exactly the locality bias the ViT lacks.
Resolution and patch size
Sequence length is N=(H/P)2, so both terms hit attention cost quadratically — and attention cost is itself quadratic in N.
| Resolution | Patch | Patches N | Sequence | Attention entries | Relative cost |
|---|---|---|---|---|---|
| 224 | 16 | 196 | 197 | 38,809 | 1× |
| 224 | 8 | 784 | 785 | 616,225 | 15.9× |
| 384 | 16 | 576 | 577 | 332,929 | 8.6× |
| 512 | 32 | 256 | 257 | 66,049 | 1.7× |
Halving the patch size from 16 to 8 costs almost 16 times more attention compute for finer detail. That trade is worth it for small objects — satellite imagery, cell nuclei, defect detection — and wasteful for whole-image classification of large subjects.
The standard recipe is to pre-train at 224 and fine-tune at 384, which improves accuracy by roughly 1–2 points on ImageNet. Remember the positional embedding interpolation when you do it.
What this means when you build something
Choose by dataset size first, not by which architecture is newer.
| Your situation | Choose | Reason |
|---|---|---|
| Under ~10k images, training from scratch | CNN (ResNet, EfficientNet) | The locality bias is free supervision you cannot afford to learn |
| Any size, fine-tuning from a pre-trained checkpoint | ViT or a hybrid | Pre-training already paid the data cost; ViT transfers exceptionally well |
| Dense prediction — segmentation, detection | Swin, or a hierarchical ViT | Plain ViT has one fixed resolution throughout; dense tasks need a multi-scale feature pyramid |
| Very high resolution or small objects | Windowed or hierarchical attention | Global attention at P=8 on a 1024-pixel image is 16,384 tokens: 2.7×108 attention entries per head |
| Multimodal, image plus text | ViT | Both modalities become token sequences that a single transformer stack can attend over jointly |
Three failure modes worth naming because each looks like something else:
"ViT is worse than my ResNet." Almost always a from-scratch run on a small dataset. Check whether you loaded pre-trained weights at all. If you did and it is still worse, check the learning rate: ViTs want AdamW at around 10−4 with weight decay 0.05 for fine-tuning, not the SGD-with-momentum recipe that works for CNNs.
"The attention maps are meaningless." Nearly always head averaging, or plotting only the final layer. Set average_attn_weights=False, look at individual heads, and use rollout for a whole-network view.
"It works at 224 but breaks at 384." The positional embedding table has the wrong number of rows. Interpolate it.
The broader lesson generalises past vision. The transformer did not win because attention is a better image operator than convolution — for a fixed, modest data budget it is not. It won because it assumes less, and an architecture that assumes less scales further when you can afford to replace assumptions with data. That is the same reason the same architecture now handles audio, video, protein sequences and point clouds with barely any modification: the tokenisation changes, the mechanism does not.