Attention Mechanisms and Transformers

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.

A 224 by 224 image becomes 197 tokensImage, 224by 224 by 3Cut into196 patchesof 16 by 16Flatten eachto 768 numbersLinearprojection,plus positionPrepend CLS,giving 197 tokensThe projection is a stride-16 convolution, which is how it is actually implemented.
Patching is the whole adaptation — everything after it is the language transformer, unchanged.

Two different assumptions about images

A convolutional layer hard-codes three beliefs about the visual world:

Built-in assumptionHow the architecture enforces itWhen it is wrong
Locality — a pixel relates most to its neighboursA 3×33\times3 kernel sees nine pixels and nothing elseRelating two objects at opposite corners of the frame
Translation equivariance — a cat is a cat wherever it appearsThe same kernel weights slide across every positionTasks where absolute position matters (document layout, medical scans with fixed anatomy)
Hierarchy — edges compose into textures, textures into parts, parts into objectsStacked layers with pooling, receptive field growing with depthRarely 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×33\times3 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

224−12≈112 layers\frac{224 - 1}{2} \approx 112 \text{ layers}

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×224224\times224 image has 50,176 pixels, so the attention matrix would be

50,1762=2.52×109 entries50{,}176^2 = 2.52 \times 10^{9} \text{ entries}

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×WH \times W and patch size PP:

N=HWP2N = \frac{HW}{P^2}

With H=W=224H = W = 224 and P=16P = 16: N=2242/162=50176/256=196N = 224^2/16^2 = 50176/256 = 196 patches, arranged 14×1414 \times 14.

Each patch is 16×16×3=76816 \times 16 \times 3 = 768 raw numbers. Flatten it and project it linearly to the model dimension DD:

zi=xp(i)E,E∈R768×Dz_i = \mathbf{x}_p^{(i)} E, \qquad E \in \mathbb{R}^{768 \times D}

For ViT-Base, D=768D = 768, so EE is 768×768768 \times 768 — 589,824 weights plus 768 biases, giving 590,592 parameters in the patch embedding.

The attention matrix is now 197×197=38,809197 \times 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 PP and stride PP:

Python
import torchimport torch.nn as nnclass PatchEmbedding(nn.Module):    def __init__(self, img_size=224, patch_size=16, in_ch=3, embed_dim=768):        super().__init__()        assert img_size % patch_size == 0, "image size must divide by patch size"        self.n_patches = (img_size // patch_size) ** 2        # kernel == stride == patch_size: each output position sees exactly one patch,        # with no overlap. Mathematically identical to flatten-then-Linear.        self.proj = nn.Conv2d(in_ch, embed_dim,                              kernel_size=patch_size, stride=patch_size)    def forward(self, x):                    # (B, 3, 224, 224)        x = self.proj(x)                     # (B, 768, 14, 14)        x = x.flatten(2)                     # (B, 768, 196)        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.

Python
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))     # 768 parameters...cls = self.cls_token.expand(B, -1, -1)          # (B, 1, 768)x = 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,296197 \times 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×1414\times14 grid, bicubically resize to the new grid, flatten back:

Python
import torch.nn.functional as Fdef interpolate_pos_embed(pos_embed, old_grid, new_grid):    """pos_embed: (1, 1 + old_grid**2, D). Keeps the CLS row, resizes the rest."""    cls_tok, patch_tok = pos_embed[:, :1], pos_embed[:, 1:]    D = patch_tok.shape[-1]    patch_tok = patch_tok.reshape(1, old_grid, old_grid, D).permute(0, 3, 1, 2)    patch_tok = F.interpolate(patch_tok, size=(new_grid, new_grid),                              mode='bicubic', align_corners=False)    patch_tok = patch_tok.permute(0, 2, 3, 1).reshape(1, new_grid ** 2, D)    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

Python
class ViTBlock(nn.Module):    """Pre-norm transformer block: LN inside the residual branch."""    def __init__(self, dim, heads, mlp_ratio=4.0, dropout=0.0):        super().__init__()        self.norm1 = nn.LayerNorm(dim)        self.attn = nn.MultiheadAttention(dim, heads, dropout=dropout,                                          batch_first=True)        self.norm2 = nn.LayerNorm(dim)        hidden = int(dim * mlp_ratio)        self.mlp = nn.Sequential(            nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(dropout),            nn.Linear(hidden, dim), nn.Dropout(dropout))    def forward(self, x, return_attn=False):        h = self.norm1(x)        a, w = self.attn(h, h, h,                         need_weights=return_attn,                         average_attn_weights=False)   # keep per-head weights        x = x + a        x = x + self.mlp(self.norm2(x))        return (x, w) if return_attn else (x, None)class VisionTransformer(nn.Module):    def __init__(self, img_size=224, patch_size=16, in_ch=3, num_classes=1000,                 dim=768, depth=12, heads=12, mlp_ratio=4.0, dropout=0.0):        super().__init__()        self.patch_embed = PatchEmbedding(img_size, patch_size, in_ch, dim)        n = self.patch_embed.n_patches        self.cls_token = nn.Parameter(torch.zeros(1, 1, dim))        self.pos_embed = nn.Parameter(torch.zeros(1, n + 1, dim))        self.pos_drop  = nn.Dropout(dropout)        self.blocks = nn.ModuleList(            ViTBlock(dim, heads, mlp_ratio, dropout) for _ in range(depth))        self.norm = nn.LayerNorm(dim)              # required with pre-norm blocks        self.head = nn.Linear(dim, num_classes)        nn.init.trunc_normal_(self.pos_embed, std=0.02)        nn.init.trunc_normal_(self.cls_token, std=0.02)    def forward(self, x, return_attn=False):        B = x.shape[0]        x = self.patch_embed(x)                                    # (B, 196, D)        x = torch.cat([self.cls_token.expand(B, -1, -1), x], 1)    # (B, 197, D)        x = self.pos_drop(x + self.pos_embed)        maps = []        for blk in self.blocks:            x, w = blk(x, return_attn)            if return_attn:                maps.append(w)        x = self.norm(x)        logits = self.head(x[:, 0])            # the CLS position only        return (logits, maps) if return_attn else logits

The 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

ModelLayersDDHeadsMLPParameters
ViT-Base1276812307286 M
ViT-Large241024164096307 M
ViT-Huge321280165120632 M

Verify ViT-Base from its parts. Per block: attention is four 768×768768\times768 matrices with biases, 4(7682+768)=2,362,3684(768^2 + 768) = 2{,}362{,}368; the MLP is 768 ⁣⋅ ⁣3072+3072+3072 ⁣⋅ ⁣768+768=4,722,432768\!\cdot\!3072 + 3072 + 3072\!\cdot\!768 + 768 = 4{,}722{,}432; two LayerNorms add 4×768=3,0724 \times 768 = 3{,}072. That is 7,087,872 per block, and 12×7,087,872=85,054,46412 \times 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×1414\times14 heatmap.

Python
import matplotlib.pyplot as plt@torch.no_grad()def cls_attention_map(model, image, layer=-1, head=None, grid=14):    model.eval()    _, maps = model(image.unsqueeze(0), return_attn=True)    w = maps[layer][0]                       # (heads, 197, 197)    w = w.mean(0) if head is None else w[head]    cls_row = w[0, 1:]                       # drop the CLS-to-CLS entry    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:

TrapWhat you seeFix
F.scaled_dot_product_attention is used internallyNo weights at all — the fused kernel never materialises the matrixRegister a forward hook on the explicit path, or set the library flag that disables fusion for inspection
Forgetting to drop index 0 before reshaping197 values will not reshape to 14×1414\times14, or you silently shift every patch by oneSlice [1:] before reshape(14, 14)
Plotting only the last layerDiffuse maps with strong sinks on a few patchesUse 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:

A~(ℓ)=12(A(ℓ)+I),R=A~(L)A~(L−1)⋯A~(1)\tilde{A}^{(\ell)} = \tfrac{1}{2}\left(A^{(\ell)} + I\right), \qquad R = \tilde{A}^{(L)}\tilde{A}^{(L-1)}\cdots\tilde{A}^{(1)}

Python
@torch.no_grad()def attention_rollout(maps):    """maps: list of (B, heads, T, T). Returns (T, T) for the first example."""    T = maps[0].shape[-1]    result = torch.eye(T)    for a in maps:        a = a[0].mean(0)                      # average heads within a layer        a = 0.5 * a + 0.5 * torch.eye(T)      # account for the residual path        a = a / a.sum(dim=-1, keepdim=True)   # renormalise rows to sum to 1        result = a @ result    return result

Rollout 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 dataBest CNN (BiT / ResNet)ViTWinner
ImageNet-1k (1.3 M)~77%~72–77% for ViT-BCNN
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)2N = (H/P)^2, so both terms hit attention cost quadratically — and attention cost is itself quadratic in NN.

ResolutionPatchPatches NNSequenceAttention entriesRelative cost
2241619619738,8091×
2248784785616,22515.9×
38416576577332,9298.6×
5123225625766,0491.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 situationChooseReason
Under ~10k images, training from scratchCNN (ResNet, EfficientNet)The locality bias is free supervision you cannot afford to learn
Any size, fine-tuning from a pre-trained checkpointViT or a hybridPre-training already paid the data cost; ViT transfers exceptionally well
Dense prediction — segmentation, detectionSwin, or a hierarchical ViTPlain ViT has one fixed resolution throughout; dense tasks need a multi-scale feature pyramid
Very high resolution or small objectsWindowed or hierarchical attentionGlobal attention at P=8P=8 on a 1024-pixel image is 16,38416{,}384 tokens: 2.7×1082.7\times10^{8} attention entries per head
Multimodal, image plus textViTBoth 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−410^{-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.