ViT, explained · Part 2 of 2 · Covers §2, §3.1, Figure 1, Eq. 1 to 4, Appendix A

Inside ViT: Patches, Embeddings and the Encoder

Section 2, Figure 1, Section 3.1 and Appendix A, line by line: the earlier attempts to put attention on images, then every step from a 224×224 picture to a class label, with Equations 1 to 8 recomputed by hand on the released ViT-B/16 and matched against the library at every stage, the 86,567,656 parameters counted exactly, and a four-token attention example small enough to check with a pencil.

An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. Alexey Dosovitskiy, Lucas Beyer, Alexander Kolesnikov, Dirk Weissenborn, Xiaohua Zhai, Thomas Unterthiner, Mostafa Dehghani, Matthias Minderer, Georg Heigold, Sylvain Gelly, Jakob Uszkoreit, Neil Houlsby. ICLR 2021, 2020. arXiv:2010.11929

In Part 1 we read the paper's claim: cut a picture into 16×16 patches, feed them to an unchanged Transformer, and with enough data it beats the best convolutional networks. This part opens the model. First we read Section 2, the earlier work the paper builds on and argues against. Then Figure 1 and Section 3.1, which describe the model in one picture, four equations and a few paragraphs. Last we read Appendix A, the attention equations the paper hands to its appendix.

Everything that can be checked, we check on a released model, google/vit-base-patch16-224 (ViT-B/16, pre-trained on ImageNet-21k and fine-tuned on ImageNet at 224 pixels), with one real picture of two cats. Every step of the model is recomputed in plain PyTorch and compared with the library's own output.

Two of these ideas are worth seeing in their own papers. Here is Parmar et al. (2018), the Image Transformer, saying in its abstract why it went local:

And here are the patterns of the Sparse Transformer, which the ViT paper names explicitly:

The Cordonnier paper is about something else: it proves that attention can do what a convolution does, and shows that trained attention layers often learn to. Its 2×2 step appears in the experiments as a down-sampling detail:

Full self-attention over pixels: cost grows with (number of pixels)²Local attentionParmar 2018 (ImageTransformer): eachpixel attends to aneighbourhood.Hu, Ramachandran,Zhao: replace convsSparse and axialChild 2019 (SparseTransformers);Weissenborn 2019(blocks); Ho 2019,Wang 2020a: oneaxis at a timeSmall patchesCordonnier 2020:2×2 patches, fullattention on top.Small images only.Closest to ViT.CNN + attentionBello 2019; Hu 2018;Carion 2020 (DETR);Wang 2018; Sun 2019;Wu 2020; Locatello2020; UNITER,ViLBERT, VisualBERTPixels as tokensiGPT (Chen 2020a):generative modelon down-scaledpixels; linearprobe reaches 72%on ImageNetViT: 16×16 patches, full (global) attention, the standard encoder unchanged,and pre-training on 14M to 300M imagesthe large-data line it joins: Mahajan 2018, Touvron 2019, Xie 2020, Sun 2017, Kolesnikov 2020 (BiT), Djolonga 2020
The related work on one map. Five earlier ways of putting attention on images, each a different answer to the quadratic cost: local attention, sparse and axial patterns, 2×2 patches with full attention, attention bolted on to a CNN, and iGPT on shrunken pixels. ViT, below, keeps full attention and the standard encoder, uses 16×16 patches, and joins the large-data line of work.

Figure 1: the whole model in one picture

Section 3 is short. It opens with one sentence of design philosophy and one picture, Figure 1. Here is the picture first, because everything else in this part explains one piece of it.

How to read Figure 1, left to right and bottom to top:

  • Bottom left, the picture: split into a grid of patches (nine in the drawing; 196 for a real 224×224 picture with 16×16 patches).
  • "Linear Projection of Flattened Patches" (pink bar): each patch is flattened into one long row of numbers and multiplied by one matrix, E\mathbf{E}. Every patch uses the same matrix.
  • "Patch + Position Embedding" (the circles 0 to 9): to each projected patch a learned position vector is added, so the model can tell patch 3 from patch 7. The starred circle 0 is the "extra learnable [class] embedding": a token that belongs to no patch.
  • "Transformer Encoder" (grey box): the standard encoder, LL layers of it. Every token can look at every token.
  • "MLP Head" and "Class": only the output at position 0, the class token, is read. A small head turns it into a label.
  • Right half, the encoder block: "Norm", "Multi-Head Attention", an addition; "Norm", "MLP", an addition. Note that Norm comes before each block, with the plus signs after. That is the "pre-norm" arrangement, and it is Equations 2 and 3.
1The picture224 × 224 pixels,3 colour channelsx: H × W × C2Cut into patches16 × 16 × 3 each224/16 = 14 per sideN = 14 × 14 = 1963Flatten each patch768numbersper patch4Linear projection EE768×768x_p E: 196 patch embeddings of 7685Prepend [class][class]p1p2p3…p196196 + 1 = 197 tokens; x_class is learned6Add position embeddings[class]p1p2p3…p196pos 0pos 1pos 2pos 3…pos 196+E_pos: 197 learned rows (Eq. 1)7Transformer encoderLN → MSA → + → LN → MLP → +repeated L = 12 times (Eq. 2, 3)197 × 768 in, 197 × 768 out8Read the [class] outputz_L⁰z_L¹z_L²…z_L¹⁹⁶y = LN(z_L⁰): 768 numbers (Eq. 4)9Classification headLinear 768→1000Egyptian cat0.937 (softmax)one linear layer when fine-tuned
Figure 1 redrawn as nine numbered frames with the real sizes. The 224×224×3 picture is cut into 196 patches of 16×16×3; each is flattened to 768 numbers and multiplied by E (768×768); the [class] token is prepended and 197 position rows are added; the 12-layer encoder keeps the shape 197×768; the [class] output is normalised and a single linear layer gives 1000 scores, the largest of which is "Egyptian cat" with probability 0.937 on our picture.

The phrase "out of the box" is also why we can check this paper so thoroughly. ViT-B/16 has the same layer count, hidden size, head count and MLP size as BERT-base (BERT Part 2), and the released code is a few hundred lines on top of a standard Transformer implementation.

BERT (text)ViT (pictures)[CLS]thekidsmiles[SEP]WordPiece tokens[class]p1p2…p19616 × 16 patchestoken table lookup + segment + position (then LayerNorm)x_p E (linear projection) + positionTransformer encoder12 layers, width 768, 12 heads, MLP 3072Transformer encoder12 layers, width 768, 12 heads, MLP 3072C = output of [CLS] → task head (sentence label)y = LN(z_L⁰) → head (image class)=≈Differences: BERT normalises after each block (post-norm) and looks tokens up in a table of 30,522 rows;ViT normalises before each block (pre-norm) and projects every patch with one 768 × 768 matrix E.
BERT and ViT side by side. BERT reads word pieces with a [CLS] token in front; ViT reads 16×16 patches with a [class] token in front. BERT embeds a token by table lookup plus segment plus position; ViT embeds a patch by a linear projection plus position. The encoder in between has the same dimensions (12 layers, width 768, 12 heads, MLP 3072), and in both the first output token feeds the task head.

From pixels to a sequence: the reshape

The arithmetic for our picture, with the paper's symbols:

N=HWP2=224×224162=50,176256=196,P2⋅C=16×16×3=768N = \frac{HW}{P^2} = \frac{224 \times 224}{16^2} = \frac{50{,}176}{256} = 196, \qquad P^2 \cdot C = 16 \times 16 \times 3 = 768

where H=W=224H = W = 224 is the resolution the model was trained at, P=16P = 16 is the patch size (the "/16" in "ViT-B/16"), and C=3C = 3 for red, green and blue. So the picture becomes xp\mathbf{x}_p, a matrix of 196 rows and 768 columns. A happy coincidence of ViT-B/16 at 224: the row length P2C=768P^2 C = 768 equals D=768D = 768, so E\mathbf{E} happens to be square. For ViT-B/32 the rows are 32⋅32⋅3=3,07232 \cdot 32 \cdot 3 = 3{,}072 long and E\mathbf{E} is 3,072×7683{,}072 \times 768.

x ∈ R^(H×W×C): 224 × 224 × 3224 px = 14 patches × 16 px14reshapeN = HW / P²= 224 · 224 / 16²= 50,176 / 256= 196 patchesx_p196 rows × 768 columns (N × P²·C)row 0 = top-left patchrow 195 = bottom-righteach row: the 768pixel numbers ofone patch (16·16·3)The grid is read row by row: patch n sits at row ⌊n/14⌋, column n mod 14.The order is fixed, so position 1 is always the top-left corner.
The reshape of Section 3.1 drawn out. Left: the 224×224×3 picture as a 14×14 grid of 16-pixel patches (224/16 = 14 per side). Right: the matrix x_p with 196 rows and 768 columns; row 0 is the top-left patch, row 195 the bottom-right. N = 224·224/16² = 50,176/256 = 196.

Here is the picture we use, resized to 224×224 the way the released preprocessing does it, with the 14×14 patch grid drawn on top:

The sample picture: two cats lying on a red sofa with remote controls, resized to 224 by 224 pixels and shown with a 14 by 14 grid of white lines marking the 16 by 16 pixel patches
Our test picture (from the `huggingface/cats-image` dataset) at 224×224 with the 196 patches marked. Each cell is one token for the model.

Cutting and flattening in PyTorch is three lines (vit_part2_math.py). unfold slides a 16-wide window with step 16 along the height and then the width, which gives one 16×16 square per patch and per channel; a permute and a reshape put the patches as rows:

python
patches = x.unfold(1, P, P).unfold(2, P, P)                      # (3, 14, 14, 16, 16)
patches = patches.permute(1, 2, 0, 3, 4).reshape(N, C * P * P)   # (196, 768): row n = patch n
plain text
== 1. cut x (3, 224, 224) into 196 patches of 16x16x3 and flatten each to 768 numbers ==
  x.unfold -> (3, 14, 14, 16, 16)  (channels, 14 rows of patches, 14 columns, 16, 16)
  x_p = flattened patches: (196, 768)   (N = 196 patches, P^2*C = 768 numbers each)
  patch 0 (top-left corner), first 8 of its 768 numbers: [  0.114,   0.169,   0.184,   0.200,   0.208,   0.239,   0.231,   0.200]
  patch 0, number 0 is red channel, pixel (0,0): x[0,0,0] = 0.114; number 256 is green (0,0): -0.804; number 512 is blue (0,0): -0.545
  check: patch 17 (row 1, col 3) == x[:, 16:32, 48:64] flattened: True

The numbers are pixel values after the released preprocessing: each channel is scaled from 0 to 255 into [0,1][0, 1] and then mapped to [−1,1][-1, 1] by (v−0.5)/0.5(v - 0.5)/0.5. The top-left pixel of our picture is reddish (red 0.114, green −0.804, blue −0.545), which is the sofa. The check on patch 17 confirms the order: patch nn sits at grid row ⌊n/14⌋\lfloor n/14 \rfloor and column n mod 14n \bmod 14, so patch 17 is row 1, column 3, and it is exactly the pixels x[:, 16:32, 48:64].

One 16 × 16 × 3 patch becomes one row of 768 numberspatch 0 (top-left)C = 3 planes of 16×16x_p[0][0] = 0.114x_p[0][1] = 0.169x_p[0][2] = 0.184x_p[0][3] = 0.200x_p[0][4] = 0.208x_p[0][5] = 0.239x_p[0][6] = 0.231x_p[0][7] = 0.200x_p[0][8] = 0.216x_p[0][9] = 0.239x_p[0][10] = 0.255x_p[0][11] = 0.286x_p[0][12] = 0.294x_p[0][13] = 0.302x_p[0][14] = 0.302x_p[0][15] = 0.294x_p[0][16] = 0.137x_p[0][17] = 0.169x_p[0][18] = 0.184x_p[0][19] = 0.224x_p[0][20] = 0.208x_p[0][21] = 0.239x_p[0][22] = 0.247x_p[0][23] = 0.224x_p[0][24] = 0.231x_p[0][25] = 0.255x_p[0][26] = 0.271x_p[0][27] = 0.294x_p[0][28] = 0.294x_p[0][29] = 0.326x_p[0][30] = 0.341x_p[0][31] = 0.326x_p[0][32] = 0.114x_p[0][33] = 0.153x_p[0][34] = 0.161x_p[0][35] = 0.176x_p[0][36] = 0.192x_p[0][37] = 0.200x_p[0][38] = 0.216x_p[0][39] = 0.216first 40 of the 768 numbers (red channel: pixel row 0, start of row 1):0.114, 0.169, 0.184, 0.200, 0.208, 0.239, 0.231, 0.200, …256 red: 16 rows × 16 pixels256 green256 blueall 768 numbers, in the order (channel, row, column):number 0 = red at pixel (0,0) = 0.114; number 256 = green at (0,0) = -0.804; number 512 = blue at (0,0) = -0.545
One 16×16×3 patch becomes one row of 768 numbers: the 256 red values first, then 256 green, then 256 blue. The strip shows the real first 40 numbers of patch 0 (the red channel of pixel row 0, then the start of row 1); the darker the square, the larger the value.

The projection E\mathbf{E}. Each row of xp\mathbf{x}_p is multiplied by the same matrix E∈R(P2⋅C)×D\mathbf{E} \in \mathbb{R}^{(P^2 \cdot C) \times D}. One matrix product does all 196 patches at once:

xp⏟196×768⋅E⏟768×768=xpE⏟196×768\underbrace{\mathbf{x}_p}_{196 \times 768} \cdot \underbrace{\mathbf{E}}_{768 \times 768} = \underbrace{\mathbf{x}_p \mathbf{E}}_{196 \times 768}
The patch embedding is one matrix multiplicationx_p196 × 768·E768 × 768=x_p E196 × 768Each row of x_p (onepatch, 768 numbers)times E gives one rowof 768 embeddingnumbers.E: 768 × 768 = 589,824weights + 768 biases.The embedding row for patch 0 starts 0.058, -0.024, -0.248, 3.689 … (768 numbers in all).D = 768 equals P²·C = 768 only by coincidence of ViT-B/16 at 224: for ViT-B/32 the rows are 32·32·3 = 3,072 longand E is 3,072 × 768.
The patch embedding is one matrix multiplication. x_p (196×768) times E (768×768) gives 196 patch embeddings of 768 numbers. E holds 768 × 768 = 589,824 weights plus a bias of 768, and it is the only thing the model learns about pixels before the Transformer.

There is one small gap between the paper and the released code here, and it is worth closing. The paper says "linear projection". The released PyTorch model (and the authors' original code) implement it as a convolution with a 16×16 kernel and a stride of 16. These are the same operation: a convolution with kernel size equal to its stride touches each patch exactly once, and what it does to a patch is a dot product with each of its 768 filters, which is a linear layer. I checked by reshaping the convolution's weight into a 768×768768 \times 768 matrix and multiplying the flattened patches with it:

python
conv = model.vit.embeddings.patch_embeddings.projection    # Conv2d(3, 768, kernel_size=16, stride=16)
E = conv.weight.reshape(768, 768).T                        # (P²·C) × D: column d is filter d, flattened
ours = patches @ E + conv.bias                             # Eq. 1's x_p E, for all 196 patches
theirs = conv(pixel_values).flatten(2).transpose(1, 2)[0]  # what the library computes
print((ours - theirs).abs().max())
plain text
== 2. the patch embedding E: the library's Conv2d(3, 768, kernel 16, stride 16) equals Linear(768 -> 768) on x_p ==
  conv weight (768, 3, 16, 16) reshaped to E (768, 768) (rows = the 768 pixel numbers of a patch, columns = D = 768)
  x_p E: (196 x 768) . (768 x 768) = (196, 768)
  max |difference|, our x_p E + b vs the library conv: 5.13e-06
  patch 0 embedding, first 6 of 768: [  0.058,  -0.024,  -0.248,   3.689,   0.356,   0.098]

A difference of 5×10−65 \times 10^{-6} is rounding noise in 32-bit arithmetic. So "linear projection of flattened patches" and "16×16 convolution with stride 16" are two names for the same 589,824 weights. The convolution is simply the faster way to run it.

1Library: Conv2d, kernel 16, stride 16a 16×16×3 filterplaced on each patch(stride 16: no overlap)768 filters →768 numbers per patchweight shape (768, 3, 16, 16)output (768, 14, 14) → (196, 768)2Eq. 1: a Linear layer on x_px_p row: 768 numbersE768×768E = conv.weight.reshape(768, 768).Tx_p @ E + biasSame 589,824 weights, just reshaped.max |difference|: 5.1e-06
Two views of the same patch embedding. Left: the library's convolution places a 16×16×3 filter on each patch with stride 16, so there is no overlap, and 768 filters give 768 numbers per patch. Right: Eq. 1's linear layer on the flattened rows. The weights are the same numbers reshaped; the outputs agree to 5e-6.

The [class] token and the head

Here is the BERT sentence the paper points to:

The two heads can be seen in the released checkpoints. The fine-tuned model google/vit-base-patch16-224 ends in exactly one linear layer, and the pre-trained-only model google/vit-base-patch16-224-in21k ends in a dense layer with a tanh (a "pooler", the hidden layer of the pre-training MLP head; the 21,843-class output layer on top of it was not released):

plain text
  head of the fine-tuned checkpoint: Linear(in_features=768, out_features=1000, bias=True)   (one linear layer, as Section 3.1 says for fine-tuning)
  the pre-trained-only checkpoint google/vit-base-patch16-224-in21k ends in: ViTPooler(
  (dense): Linear(in_features=768, out_features=768, bias=True)
  (activation): Tanh()
)
  fine-tuned checkpoint has a pooler: False;  in21k checkpoint has a pooler: True
  (the paper: an MLP with one hidden layer at pre-training time, a single linear layer at fine-tuning time)

So the released files match the sentence: an MLP with one hidden layer for pre-training, one linear layer after fine-tuning. Part 3 covers how the head is swapped and why it starts from zeros.

Position embeddings

Why are they needed at all? Because self-attention treats its input as a set. Shuffle the input rows and the output rows shuffle in exactly the same way, with the same numbers. We check this on the tiny four-token example of the Appendix A section below (vit_part2_attn.py):

plain text
== permutation equivariance: shuffle the input rows, the output rows shuffle the same way ==
  new order of the tokens: ['t2', 't0', 't3', 't1']
  MSA(z shuffled)  shape (4, 6)      (MSA(z) in the original order is printed in the Appendix A section)
       t2  1.784  2.885  0.000 -0.442 -2.007 -1.173
       t0  1.510  0.045  0.000  0.070 -1.023 -0.240
       t3  2.091  1.299  0.000 -0.228 -1.766 -0.560
       t1  1.679  0.231  0.000 -0.155 -1.116 -0.110
  MSA(z shuffled) == MSA(z) with its rows shuffled the same way: True   (max |difference| 4.8e-07)
  so without position embeddings the model cannot tell where a patch came from: only the set of patches matters.
  with a position row added before attention (z + E_pos), the shuffled run no longer matches: max |difference| 3.875

Compare with the original-order output in the Appendix A section: the rows for t2 and t0 swap places and nothing else changes. A model built only from attention and per-token MLPs would give the same class to a picture and to the same picture with its patches scrambled. Adding a different vector to each position breaks the tie: the shuffled run now differs by up to 3.875.

Reorder the input rows: the output rows reorder the same way (no positions)original ordert0MSA[1.51, 0.04, 0.00, …]t1MSA[1.68, 0.23, 0.00, …]t2MSA[1.78, 2.88, 0.00, …]t3MSA[2.09, 1.30, 0.00, …]shuffled: t2, t0, t3, t1t2MSA[1.78, 2.88, 0.00, …]t0MSA[1.51, 0.04, 0.00, …]t3MSA[2.09, 1.30, 0.00, …]t1MSA[1.68, 0.23, 0.00, …]Identical numbers, only the rows moved (max |difference| under 1e-6): attention only sees a set of tokens.Add a position row to each token before attention and the shuffled run differs (max |difference| 3.875):now the order matters.
Permutation equivariance in the tiny example. The four tokens go through multi-head self-attention in the original order (left) and in a shuffled order (right); the dashed lines join equal output rows. The numbers are identical, only the order moved. With a position row added to each token first, the two runs differ.

The encoder: LayerNorm, attention, MLP

One encoder layer: LayerNorm before each block, residual after it (Eq. 2, 3)z_{l-1}LNMSA+z'_lLNMLP+z_lresidual: + z_{l-1}residual: + z'_lEq. 2: z'_l = MSA(LN(z_{l-1})) + z_{l-1}Eq. 3: z_l = MLP(LN(z'_l)) + z'_lEvery box keeps the shape 197 × 768. LN and the MLP work on each token on its own;only MSA mixes information between tokens.
One encoder layer, left to right. The input from the previous layer goes through LayerNorm and multi-head self-attention and is added back to itself (Eq. 2, the first residual arc), then through LayerNorm and the MLP and is added back again (Eq. 3, the second arc). Every box keeps the shape 197×768. LayerNorm and the MLP work on one token at a time; only MSA mixes tokens.

Equations 1 to 4, with the real numbers

The four equations are the whole model. We take them in order and, for each, run the real ViT-B/16 on the cat picture by hand.

Equation 1: the input sequence

The idea in one sentence: put the class token in front of the 196 projected patches, then add a position vector to each of the 197.

z0=[xclass; xp1E; xp2E; ⋯ ; xpNE]+Epos,E∈R(P2⋅C)×D, Epos∈R(N+1)×D\mathbf{z}_0 = [\mathbf{x}_{\text{class}};\ \mathbf{x}_p^1 \mathbf{E};\ \mathbf{x}_p^2 \mathbf{E};\ \cdots;\ \mathbf{x}_p^N \mathbf{E}] + \mathbf{E}_{pos}, \qquad \mathbf{E} \in \mathbb{R}^{(P^2 \cdot C) \times D},\ \mathbf{E}_{pos} \in \mathbb{R}^{(N+1) \times D}

where:

  • xpn\mathbf{x}_p^n is row nn of the patch matrix, the 768 pixel numbers of patch nn (1×7681 \times 768);
  • E\mathbf{E} is the projection, 768×768768 \times 768 here (the released model also adds a bias of 768, which the equation leaves out);
  • xpnE\mathbf{x}_p^n \mathbf{E} is the patch embedding of patch nn, 1×7681 \times 768;
  • xclass\mathbf{x}_{\text{class}} is the learned class token, 1×7681 \times 768;
  • the square brackets stack the 197 rows into a 197×768197 \times 768 matrix;
  • Epos\mathbf{E}_{pos} is the table of 197 learned position vectors, 197×768197 \times 768, added row by row: row 0 to the class token, row nn to patch nn;
  • z0\mathbf{z}_0 is the result, 197×768197 \times 768, the input of layer 1.

In code, with the library's own tables and our own patch embeddings from above:

python
emb = model.vit.embeddings
x_class = emb.cls_token.detach()[0]                 # (1, 768)
E_pos = emb.position_embeddings.detach()[0]         # (197, 768)
z0 = torch.cat([x_class, our_patch_emb], 0) + E_pos # (197, 768)   Eq. 1
print((z0 - emb(pixel_values)[0]).abs().max())      # against the library's embedding layer
plain text
== 3. Eq. 1: z_0 = [x_class; x_p^1 E; ...; x_p^N E] + E_pos ==
  x_class (1, 768)   [x_class; patches] (197, 768)   E_pos (197, 768)   z_0 (197, 768)
  x_class, first 6:        [  0.010,   0.015,  -0.267,  -0.001,   0.405,   0.054]
  E_pos[0] (for [class]):  [  0.010,   0.015,  -0.267,  -0.000,   0.406,   0.054]
  z_0[0] = x_class + E_pos[0]: [  0.020,   0.030,  -0.534,  -0.001,   0.811,   0.108]
  E_pos[1] (for patch 1):  [  0.156,  -0.124,   0.447,   0.009,   0.428,  -0.462]
  z_0[1] = x_p^1 E + E_pos[1]: [  0.214,  -0.148,   0.198,   3.698,   0.783,  -0.365]
  max |difference|, our z_0 vs the library embedding output: 5.13e-06

Check the second row by hand: patch 1's embedding started 0.058,−0.024,−0.248,3.689,…0.058, -0.024, -0.248, 3.689, \dots (step 2 above); add Epos[1]=0.156,−0.124,0.447,0.009,…\mathbf{E}_{pos}[1] = 0.156, -0.124, 0.447, 0.009, \dots and you get 0.214,−0.148,0.198,3.6980.214, -0.148, 0.198, 3.698, which is what the script prints. The class token and its position vector are a curiosity: in this checkpoint they are almost identical lists of numbers (0.010,0.015,−0.267,…0.010, 0.015, -0.267, \dots in both), so z0[0]\mathbf{z}_0[0] is close to 2 xclass2\,\mathbf{x}_{\text{class}}. Nothing in the model requires that; the two vectors are only ever used as a sum, so training was free to split the sum between them any way it liked.

Eq. 1 with the real shapes and numbersx_classx_p¹Ex_p²Ex_p³Ex_p⁴E…x_p¹⁹⁶E[ … ]197 × 768E_pos[0]E_pos[1]E_pos[2]E_pos[3]E_pos[4]…E_pos[196]+197 × 768z₀[0]z₀[1]z₀[2]z₀[3]z₀[4]…z₀[196]=197 × 768the [class] position, first 6 of 768 numbers:x_class [ 0.010, 0.015, -0.267, -0.001, 0.405, 0.054, …]E_pos[0] [ 0.010, 0.015, -0.267, -0.000, 0.406, 0.054, …]z_0[0] [ 0.019, 0.030, -0.534, -0.001, 0.811, 0.108, …]the first patch: x_p¹E [ 0.058, -0.024, -0.248, 3.689, 0.356, 0.098, …]+ E_pos[1] [ 0.156, -0.124, 0.447, 0.009, 0.428, -0.462, …]= z_0[1] [ 0.214, -0.148, 0.198, 3.698, 0.783, -0.364, …]Our z_0 matches the library embedding output to 5.1e-06.
Equation 1 drawn with the real shapes. Top row: the learned class embedding and the 196 patch embeddings, 197 rows of 768. Middle: the 197 learned position rows. Bottom: their sum z_0. Below, the real first six numbers of the class position; our z_0 matches the library's embedding output to 5e-6.

Equation 2: LayerNorm, attention, residual

The idea in one sentence: normalise every token, let the tokens look at each other, and add the result back to what came in.

zℓ′=MSA(LN(zℓ−1))+zℓ−1,ℓ=1…L\mathbf{z}'_\ell = \text{MSA}(\text{LN}(\mathbf{z}_{\ell-1})) + \mathbf{z}_{\ell-1}, \qquad \ell = 1 \dots L

where:

  • zℓ−1\mathbf{z}_{\ell-1} is the output of the previous layer (z0\mathbf{z}_0 for the first), 197×768197 \times 768;
  • LN\text{LN} normalises each of the 197 rows on its own (formula below), shape unchanged;
  • MSA\text{MSA} is multi-head self-attention, Appendix A, 197×768197 \times 768 in and out;
  • + zℓ−1+\ \mathbf{z}_{\ell-1} is the residual connection;
  • zℓ′\mathbf{z}'_\ell is the intermediate result, 197×768197 \times 768, which Equation 3 continues; L=12L = 12 for ViT-B.

LayerNorm by hand. For one token's vector xx of DD numbers:

μ=1D∑i=1Dxi,σ=1D∑i=1D(xi−μ)2+ϵ,LN(x)i=γi xi−μσ+βi\mu = \frac{1}{D}\sum_{i=1}^{D} x_i, \qquad \sigma = \sqrt{\frac{1}{D}\sum_{i=1}^{D}(x_i - \mu)^2 + \epsilon}, \qquad \text{LN}(x)_i = \gamma_i\,\frac{x_i - \mu}{\sigma} + \beta_i

where μ\mu is the mean of the 768 numbers, σ\sigma their standard deviation (with a tiny ϵ=10−12\epsilon = 10^{-12} added so that it is never zero), and γ,β\gamma, \beta are two learned vectors of 768 numbers, the same for every token. Here it is on the class token entering layer 1:

python
v = z0[0]                                                  # the [class] token, 768 numbers
ln = model.vit.layers[0].layernorm_before
mu, var = v.mean(), ((v - v.mean()) ** 2).mean()
ours = ln.weight * (v - mu) / torch.sqrt(var + 1e-12) + ln.bias
plain text
  LayerNorm of the [class] token (768 numbers), eps = 1e-12:
    input, first 6:            [  0.020,   0.030,  -0.534,  -0.001,   0.811,   0.108]
    mean mu = 0.0034   variance = 0.0349   std = sqrt(var + eps) = 0.1869
    (x - mu)/std, first 6:     [  0.086,   0.141,  -2.875,  -0.023,   4.320,   0.561]
    gamma, first 6:            [  0.156,   0.191,   0.033,   0.119,   0.065,   0.126]
    beta, first 6:             [  0.008,   0.016,   0.091,   0.067,   0.007,  -0.034]
    gamma*(..)+beta, first 6:  [  0.022,   0.043,  -0.005,   0.064,   0.287,   0.036]
    library LayerNorm:         [  0.022,   0.043,  -0.005,   0.064,   0.287,   0.036]   max |difference| 6.0e-08

Check the fifth number: (0.811−0.0034)/0.1869=4.321(0.811 - 0.0034)/0.1869 = 4.321, then 0.065×4.321+0.007=0.2880.065 \times 4.321 + 0.007 = 0.288, against the printed 0.287 (the inputs shown are rounded). Notice how small γ\gamma is here (0.03 to 0.19): the trained model scales the normalised numbers down a lot before attention sees them.

LayerNorm on the [class] token entering layer 1, first 6 of its 768 numbers[0][1][2][3][4][5]x (input)0.0190.030-0.534-0.0010.8110.108x − μ0.0160.026-0.537-0.0040.8080.105(x − μ) / σ0.0860.141-2.875-0.0234.3200.561γ0.1560.1910.0330.1190.0650.126β0.0080.0160.0910.0670.007-0.034γ·(…) + β0.0220.043-0.0050.0640.2870.036μ = mean of all 768 = 0.0034σ = √(variance + ε) = 0.1869ε = 1e-12γ and β are learned,one value per position,the same for every token.input length 5.18,output length 1.50.Our hand computation matches the library LayerNorm to 6e-08.Rows 2 and 3 use the exact μ and σ, so they differ slightly from what the rounded row 1 would give.
LayerNorm step by step on the first six of the 768 numbers of the class token: the input, minus the mean (0.0034), divided by the standard deviation (0.1869), times the learned gamma, plus the learned beta. Orange cells are negative. Our hand computation matches the library to 6e-8.

Attention, head by head. The normalised 197×768197 \times 768 matrix is multiplied by three matrices to give queries, keys and values, each 197×768197 \times 768, which are split into 12 heads of 64 columns. For each head, every token's query is compared with every token's key, the 197 scores are divided by 64=8\sqrt{64} = 8 and turned into weights by softmax, and the token's output is the weighted sum of the 197 value rows. The 12 outputs (197×64197 \times 64 each) are placed side by side into 197×768197 \times 768 and multiplied by one more 768×768768 \times 768 matrix. Appendix A, at the end of this part, is the formal version; here is the real one, for the class token's query in head 1 of layer 1:

python
lay = model.vit.layers[0]; at = lay.attention
q = ln1 @ at.q_proj.weight.T + at.q_proj.bias          # (197, 768); the same for k and v
qh, kh, vh = (t.view(197, 12, 64).transpose(0, 1) for t in (q, k, v))   # (12, 197, 64)
scores = qh @ kh.transpose(1, 2) / 8.0                 # (12, 197, 197)
A = torch.softmax(scores, -1)                          # every row sums to 1
SA = A @ vh                                            # (12, 197, 64)
concat = SA.transpose(0, 1).reshape(197, 768)          # the 12 heads side by side
msa = concat @ at.o_proj.weight.T + at.o_proj.bias     # U_msa
z_prime = msa + z0                                     # Eq. 2
plain text
  q = LN(z_0) W_q + b_q: (197, 768)  (the same for k and v); split into 12 heads of 64: (12, 197, 64)
  head 1: q_class . k_j / sqrt(64) for all 197 tokens j: (197,); first 6 scores: [  5.125,  -0.698,  -0.880,  -0.946,  -1.040,  -1.146]
  softmax -> weights, first 6: [  0.761,   0.002,   0.002,   0.002,   0.002,   0.001]   sum of all 197 = 1.0000   max 0.7613   min 0.0005
  the 6 largest weights of the [class] query in head 1 of layer 1:
    token   0  weight 0.7613   [class] itself
    token  14  weight 0.0026   patch  13 = grid row  0, col 13
    token   1  weight 0.0023   patch   0 = grid row  0, col  0
    token 109  weight 0.0022   patch 108 = grid row  7, col 10
    token  13  weight 0.0021   patch  12 = grid row  0, col 12
    token 153  weight 0.0020   patch 152 = grid row 10, col 12
  weight on [class] itself: 0.7613;  on all 196 patches together: 0.2387
  weight of the [class] query on itself in each of the 12 heads of layer 1: 0.76 0.90 0.78 0.11 0.90 0.65 0.79 0.02 0.82 1.00 0.37 0.39
  head 8 spreads the [class] query most (self weight 0.022); its top-5 patches:
    token   0  weight 0.0219   [class] itself
    token   8  weight 0.0146   patch   7 = grid row  0, col  7
    token   3  weight 0.0122   patch   2 = grid row  0, col  2
    token   7  weight 0.0122   patch   6 = grid row  0, col  6
    token   4  weight 0.0120   patch   3 = grid row  0, col  3

Two things are worth noticing. First, the score of the class token against itself is 5.125 while the others are around −1, and after softmax that one score takes 0.761 of the weight. Head 1 of layer 1 mostly leaves the class token alone. The same is true of most heads in this layer (0.65 to 1.00 on itself); only heads 4, 8, 11 and 12 look outward. Second, when a head does look outward, as head 8 does, no patch gets much: the largest weight is 0.0146, and the five largest all sit in the top row of the picture (the red sofa behind the cats). This is layer 1; there is no "cat detector" yet. Part 5 measures how far the heads look in every layer and shows that later layers attend to the object.

The [class] query in layer 1: its 196 patch weights on the 14 × 14 gridhead 1: 0.76 on [class] itself, largest patch weight 0.002600 → 0: 0.880 → 1: 0.730 → 2: 0.680 → 3: 0.620 → 4: 0.560 → 5: 0.610 → 6: 0.610 → 7: 0.710 → 8: 0.640 → 9: 0.570 → 10: 0.580 → 11: 0.720 → 12: 0.810 → 13: 1.0011 → 0: 0.771 → 1: 0.611 → 2: 0.541 → 3: 0.471 → 4: 0.471 → 5: 0.471 → 6: 0.561 → 7: 0.531 → 8: 0.481 → 9: 0.471 → 10: 0.321 → 11: 0.351 → 12: 0.681 → 13: 0.6922 → 0: 0.602 → 1: 0.332 → 2: 0.302 → 3: 0.332 → 4: 0.252 → 5: 0.422 → 6: 0.382 → 7: 0.442 → 8: 0.282 → 9: 0.542 → 10: 0.382 → 11: 0.252 → 12: 0.332 → 13: 0.3833 → 0: 0.483 → 1: 0.683 → 2: 0.543 → 3: 0.383 → 4: 0.433 → 5: 0.473 → 6: 0.413 → 7: 0.333 → 8: 0.583 → 9: 0.253 → 10: 0.723 → 11: 0.403 → 12: 0.713 → 13: 0.3144 → 0: 0.524 → 1: 0.544 → 2: 0.284 → 3: 0.554 → 4: 0.244 → 5: 0.284 → 6: 0.494 → 7: 0.324 → 8: 0.474 → 9: 0.194 → 10: 0.314 → 11: 0.464 → 12: 0.494 → 13: 0.5155 → 0: 0.535 → 1: 0.425 → 2: 0.725 → 3: 0.525 → 4: 0.505 → 5: 0.525 → 6: 0.695 → 7: 0.585 → 8: 0.435 → 9: 0.285 → 10: 0.535 → 11: 0.675 → 12: 0.515 → 13: 0.5666 → 0: 0.606 → 1: 0.276 → 2: 0.436 → 3: 0.426 → 4: 0.356 → 5: 0.366 → 6: 0.416 → 7: 0.326 → 8: 0.516 → 9: 0.326 → 10: 0.556 → 11: 0.466 → 12: 0.476 → 13: 0.5677 → 0: 0.507 → 1: 0.437 → 2: 0.497 → 3: 0.477 → 4: 0.387 → 5: 0.377 → 6: 0.387 → 7: 0.307 → 8: 0.747 → 9: 0.207 → 10: 0.857 → 11: 0.277 → 12: 0.467 → 13: 0.5388 → 0: 0.328 → 1: 0.488 → 2: 0.408 → 3: 0.548 → 4: 0.428 → 5: 0.398 → 6: 0.498 → 7: 0.488 → 8: 0.408 → 9: 0.308 → 10: 0.288 → 11: 0.378 → 12: 0.458 → 13: 0.5499 → 0: 0.469 → 1: 0.549 → 2: 0.449 → 3: 0.489 → 4: 0.429 → 5: 0.339 → 6: 0.349 → 7: 0.389 → 8: 0.419 → 9: 0.489 → 10: 0.359 → 11: 0.329 → 12: 0.519 → 13: 0.581010 → 0: 0.5710 → 1: 0.2810 → 2: 0.4410 → 3: 0.5810 → 4: 0.3310 → 5: 0.3610 → 6: 0.4110 → 7: 0.3910 → 8: 0.4210 → 9: 0.4710 → 10: 0.3810 → 11: 0.4610 → 12: 0.7910 → 13: 0.521111 → 0: 0.5811 → 1: 0.6111 → 2: 0.4011 → 3: 0.3611 → 4: 0.3611 → 5: 0.3711 → 6: 0.3511 → 7: 0.4211 → 8: 0.3911 → 9: 0.3711 → 10: 0.3911 → 11: 0.4411 → 12: 0.4711 → 13: 0.521212 → 0: 0.5412 → 1: 0.6112 → 2: 0.4612 → 3: 0.4212 → 4: 0.3912 → 5: 0.3612 → 6: 0.4812 → 7: 0.4312 → 8: 0.4112 → 9: 0.3812 → 10: 0.4412 → 11: 0.3912 → 12: 0.4612 → 13: 0.521313 → 0: 0.7413 → 1: 0.5313 → 2: 0.5213 → 3: 0.4913 → 4: 0.4613 → 5: 0.5013 → 6: 0.5513 → 7: 0.5413 → 8: 0.5213 → 9: 0.5013 → 10: 0.5113 → 11: 0.6013 → 12: 0.6013 → 13: 0.63012345678910111213patch column 0 to 13head 8: 0.022 on itself, largest patch weight 0.014600 → 0: 0.750 → 1: 0.790 → 2: 0.830 → 3: 0.820 → 4: 0.780 → 5: 0.790 → 6: 0.830 → 7: 1.000 → 8: 0.730 → 9: 0.540 → 10: 0.270 → 11: 0.410 → 12: 0.680 → 13: 0.7311 → 0: 0.161 → 1: 0.461 → 2: 0.681 → 3: 0.521 → 4: 0.441 → 5: 0.271 → 6: 0.291 → 7: 0.531 → 8: 0.321 → 9: 0.111 → 10: 0.101 → 11: 0.111 → 12: 0.301 → 13: 0.5122 → 0: 0.242 → 1: 0.142 → 2: 0.102 → 3: 0.152 → 4: 0.102 → 5: 0.132 → 6: 0.112 → 7: 0.082 → 8: 0.072 → 9: 0.232 → 10: 0.062 → 11: 0.072 → 12: 0.132 → 13: 0.1033 → 0: 0.333 → 1: 0.173 → 2: 0.213 → 3: 0.173 → 4: 0.083 → 5: 0.133 → 6: 0.113 → 7: 0.123 → 8: 0.053 → 9: 0.083 → 10: 0.183 → 11: 0.083 → 12: 0.153 → 13: 0.1144 → 0: 0.544 → 1: 0.584 → 2: 0.074 → 3: 0.154 → 4: 0.094 → 5: 0.114 → 6: 0.394 → 7: 0.264 → 8: 0.114 → 9: 0.064 → 10: 0.094 → 11: 0.154 → 12: 0.464 → 13: 0.3755 → 0: 0.625 → 1: 0.385 → 2: 0.075 → 3: 0.115 → 4: 0.085 → 5: 0.075 → 6: 0.175 → 7: 0.295 → 8: 0.065 → 9: 0.125 → 10: 0.085 → 11: 0.275 → 12: 0.445 → 13: 0.6266 → 0: 0.606 → 1: 0.146 → 2: 0.136 → 3: 0.066 → 4: 0.236 → 5: 0.406 → 6: 0.086 → 7: 0.406 → 8: 0.126 → 9: 0.106 → 10: 0.086 → 11: 0.406 → 12: 0.396 → 13: 0.6377 → 0: 0.427 → 1: 0.177 → 2: 0.137 → 3: 0.167 → 4: 0.467 → 5: 0.477 → 6: 0.537 → 7: 0.287 → 8: 0.077 → 9: 0.057 → 10: 0.127 → 11: 0.207 → 12: 0.477 → 13: 0.6588 → 0: 0.288 → 1: 0.118 → 2: 0.098 → 3: 0.278 → 4: 0.408 → 5: 0.508 → 6: 0.608 → 7: 0.428 → 8: 0.128 → 9: 0.188 → 10: 0.108 → 11: 0.088 → 12: 0.148 → 13: 0.6099 → 0: 0.279 → 1: 0.189 → 2: 0.079 → 3: 0.139 → 4: 0.099 → 5: 0.409 → 6: 0.469 → 7: 0.539 → 8: 0.159 → 9: 0.109 → 10: 0.229 → 11: 0.099 → 12: 0.309 → 13: 0.721010 → 0: 0.3310 → 1: 0.0910 → 2: 0.0610 → 3: 0.1410 → 4: 0.3910 → 5: 0.4510 → 6: 0.4610 → 7: 0.5210 → 8: 0.1810 → 9: 0.1310 → 10: 0.4310 → 11: 0.3510 → 12: 0.1410 → 13: 0.641111 → 0: 0.4711 → 1: 0.1411 → 2: 0.4311 → 3: 0.4511 → 4: 0.4511 → 5: 0.4111 → 6: 0.4511 → 7: 0.4911 → 8: 0.5111 → 9: 0.5211 → 10: 0.4711 → 11: 0.5911 → 12: 0.5111 → 13: 0.711212 → 0: 0.5412 → 1: 0.1612 → 2: 0.4812 → 3: 0.4912 → 4: 0.4512 → 5: 0.4412 → 6: 0.5012 → 7: 0.5012 → 8: 0.5312 → 9: 0.5012 → 10: 0.5612 → 11: 0.4012 → 12: 0.6412 → 13: 0.551313 → 0: 0.7213 → 1: 0.1313 → 2: 0.6113 → 3: 0.6213 → 4: 0.5813 → 5: 0.6013 → 6: 0.6913 → 7: 0.5913 → 8: 0.6413 → 9: 0.5713 → 10: 0.6213 → 11: 0.7513 → 12: 0.8013 → 13: 0.49012345678910111213patch column 0 to 13Rows are patch rows 0 to 13. Shade = weight divided by the largest weight in that head, so each headis on its own scale. Hover a cell for the value.
The class token's 196 patch weights in layer 1, drawn on the 14×14 grid, for two heads of the real model. Head 1 (left) puts 0.76 on the class token itself and at most 0.0026 on any patch. Head 8 (right) puts only 0.022 on itself and spreads its weight, mostly along the top row of the picture. Each heatmap is scaled to its own largest weight.
Head 8 of layer 1: where the [class] token looks, on the patch grid of the picture1234the picture, 14 × 14 patchesrank weight patch (row, col) - 0.0219 [class] itself 1 0.0146 patch 7 (row 0, col 7) 2 0.0122 patch 2 (row 0, col 2) 3 0.0122 patch 6 (row 0, col 6) 4 0.0120 patch 3 (row 0, col 3)All 197 weights of this row add up to 1, so no single patchcan be large: the top patch has 0.015, and the top row of thepicture (the red sofa behind the cats) takes most of the weight.In head 1 the same row puts 0.76 on [class] itself and at most0.0026 on any patch. Part 5 looks at what later layers attend to.
Head 8 of layer 1 on top of the patch grid of the picture, with the five strongest patches outlined and numbered. All five are in the top row (patches 7, 2, 6 and 3 of row 0), and the largest weight is 0.0146: with 197 weights that add up to 1, no single patch can be large.

The weighted sum, the join of the heads, the output matrix and the residual, still for the class token:

plain text
  SA_1 for [class] = sum_j A_0j v_j: (64,), first 6: [ -0.019,   0.029,  -0.053,   0.131,  -0.007,   0.001]
  concat of 12 heads: (197, 768);  times U_msa (768, 768) + bias: (197, 768)
  Eq. 2: z'_1 = MSA + z_0, [class] first 6: [ -0.178,  -0.015,  -0.848,  -0.393,   0.661,   0.162]
Inside MSA for ViT-B/16: 197 tokens, 12 heads of 64LN(z)197 × 768head h (12 of these)q, k, v = LN(z) U_qkveach 197 × 64A = softmax(q kᵀ/√64)197 × 197, rows sum to 1SA_h = A v: 197 × 64concat197 × 768U_msa768 × 768MSA(z)197 × 76812 heads × 64 numbers = 768 = D, so the concatenation has the same width as the input.Setting D_h = D/k keeps the parameter count the same whatever the number of heads k is (Appendix A).Attention parameters per layer: four 768 × 768 matrices (q, k, v, U_msa) and four biases = 2,362,368.
The shapes inside multi-head self-attention for ViT-B/16. The normalised input (197×768) goes through 12 heads; in each, q, k and v are 197×64, the attention matrix is 197×197 with rows that sum to 1, and the output is 197×64. The 12 outputs are concatenated to 197×768 and multiplied by U_msa (768×768). Four 768×768 matrices and four biases make 2,362,368 attention parameters per layer.

Equation 3: LayerNorm, MLP, residual

The idea in one sentence: normalise again, push every token through the same small two-layer network, and add the result back.

zℓ=MLP(LN(zℓ′))+zℓ′,ℓ=1…L\mathbf{z}_\ell = \text{MLP}(\text{LN}(\mathbf{z}'_\ell)) + \mathbf{z}'_\ell, \qquad \ell = 1 \dots L

where:

  • zℓ′\mathbf{z}'_\ell is the output of Equation 2, 197×768197 \times 768;
  • LN\text{LN} is a second LayerNorm with its own γ,β\gamma, \beta;
  • MLP(x)=GELU(xW1+b1) W2+b2\text{MLP}(x) = \text{GELU}(x W_1 + b_1)\, W_2 + b_2 with W1W_1 of 768×3,072768 \times 3{,}072 and W2W_2 of 3,072×7683{,}072 \times 768, applied to each of the 197 rows separately;
  • zℓ\mathbf{z}_\ell is the output of layer ℓ\ell, 197×768197 \times 768, the input of the next layer.

And GELU itself, number by number:

GELU(x)=x⋅Φ(x),Φ(x)=12(1+erf⁡ ⁣(x/2))\text{GELU}(x) = x \cdot \Phi(x), \qquad \Phi(x) = \tfrac{1}{2}\left(1 + \operatorname{erf}\!\left(x/\sqrt{2}\right)\right)

where Φ(x)\Phi(x) is the probability that a standard normal random number is below xx: close to 0 for very negative xx, 0.5 at 0, close to 1 for large xx. So GELU keeps a number roughly in proportion to how "positive" it is. The script computes it from torch.erf and compares with the library's F.gelu:

plain text
  GELU(x) = x * Phi(x) at a few values, against torch.nn.functional.gelu:
    x:                 -3.00   -2.00   -1.00   -0.50    0.00    0.50    1.00    2.00    3.00
    Phi(x):           0.0013  0.0228  0.1587  0.3085  0.5000  0.6915  0.8413  0.9772  0.9987
    x*Phi(x):        -0.0040 -0.0455 -0.1587 -0.1543  0.0000  0.3457  0.8413  1.9545  2.9960
    F.gelu(x):       -0.0040 -0.0455 -0.1587 -0.1543  0.0000  0.3457  0.8413  1.9545  2.9960
    relu(x):          0.0000  0.0000  0.0000  0.0000  0.0000  0.5000  1.0000  2.0000  3.0000

GELU(−1)=−1×0.1587=−0.159\text{GELU}(-1) = -1 \times 0.1587 = -0.159, not 0 as ReLU would give; GELU(2)=2×0.9772=1.955\text{GELU}(2) = 2 \times 0.9772 = 1.955, almost the 2 of ReLU. The two agree exactly with the library's function. Now the MLP on the class token in layer 1:

python
ln2 = lay.layernorm_after(z_prime)                               # second LayerNorm, own gamma and beta
h = ln2 @ lay.mlp.fc1.weight.T + lay.mlp.fc1.bias                # (197, 3072)
mlp = gelu_exact(h) @ lay.mlp.fc2.weight.T + lay.mlp.fc2.bias    # (197, 768)
z1 = mlp + z_prime                                               # Eq. 3
plain text
  LN(z'_1) [class] first 6:   [ -0.055,  -0.043,  -0.247,   0.003,   0.249,   0.020]
  MLP: LN(z') W_1 + b_1 -> (197, 3072); GELU; W_2 + b_2 -> (197, 768)
  [class] before GELU, first 6: [ -2.016,  -2.504,   0.053,  -1.484,  -0.649,   0.060]
  [class] after GELU, first 6:  [ -0.044,  -0.015,   0.028,  -0.102,  -0.168,   0.031]
  share of the 3072 numbers that are positive before GELU ([class]): 0.071
  Eq. 3: z_1 = MLP + z'_1, [class] first 6: [ -0.212,  -0.046,  -0.840,  -0.748,   0.774,   0.083]
  max |difference|, our z_1 vs the library hidden_states[1]: 6.68e-06
  max |difference|, our head-1 weights vs the library attentions[0] (all 12 heads): 2.68e-06

Only 7.1% of the 3,072 hidden numbers of the class token are positive before GELU in this layer; most of the wide hidden layer is "off" for this token, and the off numbers come out small but not zero (−2.016→−0.044-2.016 \to -0.044). Our z1\mathbf{z}_1 matches the library's output of layer 1 to 6.7×10−66.7 \times 10^{-6}, and our attention weights match its weights to 2.7×10−62.7 \times 10^{-6}. So Equations 2 and 3, as written above, are the complete layer.

The MLP block: 768 → 3072 → 768 with GELU in betweenLN(z′): 768W₁ (768 × 3072) + b₁h: 3072 numbersGELU, number by numberGELU(h): 3072W₂ (3072 × 768) + b₂MLP(…): 768-4-2241234GELU(-3.0) = -0.0040GELU(-2.0) = -0.0455GELU(-1.0) = -0.1587GELU(-0.5) = -0.1543GELU(0.0) = 0.0000GELU(0.5) = 0.3457GELU(1.0) = 0.8413GELU(2.0) = 1.9545GELU(3.0) = 2.9960GELU(x) = x · Φ(x)dashed: ReLU = max(0, x)GELU(−1) = −0.159GELU(2) = 1.955In layer 1 only 7.1% of the 3072 numbers of the [class] token are positive before GELU;the negative ones come out small but not zero.MLP parameters per layer: 768·3072 + 3072 + 3072·768 + 768 = 4,722,432 (8D² + 5D).
Left: the MLP block of one layer, 768 numbers widened to 3,072 by W₁, passed through GELU and narrowed back to 768 by W₂; 4,722,432 parameters per layer, two thirds of the layer. Right: the GELU curve x·Φ(x) computed in code, with ReLU dashed; GELU(−1) = −0.159 and GELU(2) = 1.955.

All twelve layers. The same two equations run 12 times, each layer with its own weights. Our loop against the library, layer by layer:

plain text
== 5. all 12 layers, our loop vs the library ==
  layer  1: z_1 (197, 768)   max |difference| vs hidden_states[1]: 6.68e-06   (largest value in z_1:    21.2, so relative 3.1e-07)
  layer  2: z_2 (197, 768)   max |difference| vs hidden_states[2]: 8.34e-06   (largest value in z_2:    23.7, so relative 3.5e-07)
  layer  3: z_3 (197, 768)   max |difference| vs hidden_states[3]: 9.54e-06   (largest value in z_3:    24.1, so relative 4.0e-07)
  layer  4: z_4 (197, 768)   max |difference| vs hidden_states[4]: 8.39e-05   (largest value in z_4:    56.0, so relative 1.5e-06)
  layer  5: z_5 (197, 768)   max |difference| vs hidden_states[5]: 9.16e-04   (largest value in z_5:   336.9, so relative 2.7e-06)
  layer  6: z_6 (197, 768)   max |difference| vs hidden_states[6]: 1.50e-03   (largest value in z_6:   989.2, so relative 1.5e-06)
  ... (layers 7 to 11: 1.65e-03 to 1.77e-03, relative 1.0e-06 to 1.1e-06)
  layer 12: z_12 (197, 768)   max |difference| vs hidden_states[12]: 1.83e-03   (largest value in z_12:  1699.6, so relative 1.1e-06)
  the final logits from our z_12 vs the library differ by at most 2.86e-06 (see step 6)

An honesty note on the numbers in the middle column: the absolute difference grows from 7×10−67 \times 10^{-6} in layer 1 to 1.8×10−31.8 \times 10^{-3} in layer 12. That is not a bug in the equations. The residual stream of a trained ViT contains a few enormous values (the largest number in z12\mathbf{z}_{12} is about 1,700, against 21 in z1\mathbf{z}_1), and 32-bit floating point keeps about seven significant digits, so the rounding noise grows with the values. Relative to the largest value the difference stays at about 10−610^{-6} in every layer, and the final class scores agree to 3×10−63 \times 10^{-6}.

Equation 4: the image representation

The idea in one sentence: after the last layer, normalise the class token one more time; that vector is the picture.

y=LN(zL0)\mathbf{y} = \text{LN}(\mathbf{z}_L^0)

where zL0\mathbf{z}_L^0 is row 0 (the class token) of the last layer's output, 768 numbers, and LN\text{LN} is a final LayerNorm with its own γ,β\gamma, \beta. y\mathbf{y}, also 768 numbers, is the image representation. The 196 patch rows of zL\mathbf{z}_L are computed and then ignored. The head is not in the equation; in the fine-tuned model it is logits=y Whead+bhead\text{logits} = \mathbf{y}\,W_{\text{head}} + b_{\text{head}} with WheadW_{\text{head}} of 768×1000768 \times 1000:

python
y = model.vit.layernorm(zL[0])                                 # Eq. 4, 768 numbers
logits = y @ model.classifier.weight.T + model.classifier.bias # (1000,)
probs = torch.softmax(logits, -1)
plain text
== 6. Eq. 4: y = LN(z_L^0), then the classification head ==
  z_L^0 (the [class] token after 12 layers), first 6: [  2.313,   5.512,  11.788,   0.577,   6.548,  -2.913]   mean 0.073 var 34.490
  y = LN(z_L^0), first 6: [  0.294,   0.835,   1.904,   0.081,   1.039,  -0.518]   max |difference| vs the library: 3.81e-06
  head of the fine-tuned checkpoint: Linear(in_features=768, out_features=1000, bias=True)   (one linear layer, as Section 3.1 says for fine-tuning)
  logits = y W_head + b: (1000,)   max |difference| vs model.logits: 2.86e-06
  top-5 ImageNet classes:
     0.937  class 285  Egyptian cat
     0.038  class 281  tabby, tabby cat
     0.014  class 282  tiger cat
     0.003  class 287  lynx, catamount
     0.001  class 284  Siamese cat, Siamese

The model says "Egyptian cat" with probability 0.937, and the next four guesses are all cats. Everything from the 50,176 pixels to this answer was computed above with nothing but matrix multiplications, LayerNorm, softmax and GELU, and matched the library at every step.

Inductive bias

The paragraph's two claims about ViT are easy to confirm from what we computed. The MLP is "local and translationally equivariant" because it acts on one token at a time with the same W1,W2W_1, W_2 for every token; the attention is "global" because every row of the 197×197197 \times 197 matrix had a non-zero weight for every column (the smallest weight in the class token's row was 0.0005, not 0). And the position embeddings carry no 2D information at the start because they are 197 independent learnable rows; nothing in Equation 1 says that row 15 is below row 1.

Where the 2D structure of the picture enters ViT: only at two pointscut into 16 × 16 patches2D used here:neighbouring pixelsstay in one patchE + learned positionsno 2D at the start:E_pos rows are randomuntil trainedencoder:global MSA, per-token MLPno 2D: every patch seesevery patch; the MLPtreats each token alikefine-tune at a new size:2D interpolation of E_pos2D used here:position rows laid onthe grid and resizedCNN: locality, 2D neighbourhoods and translation equivariance are built into every layer.ViT: the MLP is local (one token at a time) and translation equivariant (the same weights for every token);self-attention is global. Which patch is next to which has to be learned from data, which is why ViTneeds so much of it (Part 4).
The two places where the 2D structure of the picture enters ViT, highlighted: cutting into 16×16 patches at the start, and the 2D interpolation of position embeddings when fine-tuning at a new resolution (Part 3). The embedding and the encoder in between carry no built-in knowledge of the grid; the position rows start random and the global attention treats all patches alike.

Hybrid architecture

In shapes: a ResNet turns the 224×224×3224 \times 224 \times 3 picture into a 14×14×1,02414 \times 14 \times 1{,}024 feature map (with the extended-stage-3 option the paper pairs with ViT-B/16). Reading each position as a 1×11 \times 1 "patch" gives N=196N = 196 tokens of 1,0241{,}024 numbers, so E\mathbf{E} becomes 1,024×7681{,}024 \times 768 and everything from Equation 1 on is unchanged. The sequence length is the same 196 as ViT-B/16, which is why the paper pairs them.

Hybrid: the patches come from a CNN feature map instead of raw pixelspicture224 × 224 × 3ResNet stages(convolutions)feature map14 × 14 × 10241 × 1 "patches":196 vectorsof 1024 numbersE1024 → 768[class] + 196 tokens+ positions, thenthe same encoderThe paper (Section 4.1) takes the 7 × 7 output of stage 4 of a ResNet50, or the 14 × 14 output of astage 3 extended to replace stage 4, and feeds every position of the map as one token.Everything after E is unchanged. The 14 × 14 map gives 196 tokens, as in ViT-B/16.
The hybrid architecture. The picture goes through ResNet stages; the 14×14×1024 feature map is read as 196 one-by-one patches of 1,024 numbers; E (now 1,024 × 768) projects them; the class token and position embeddings are added; and the same encoder follows. Everything after E is identical to the plain ViT.

Table 1, as far as this part needs it

The next part reads the experimental setup in full. We need one table from it now, because every number above (768, 12, 3,072) came from it:

Table 1 drawn: one block per layer, block width ∝ hidden size DViT-BaseL = 12 layers, D = 768MLP 3072, 12 headspaper: 86M; ours: 85.8MViT-LargeL = 24 layers, D = 1024MLP 4096, 16 headspaper: 307M; ours: 303.3MViT-HugeL = 32 layers, D = 1280MLP 5120, 16 headspaper: 632M; ours: 630.8M"ours" counts the model without any classification head (the formula of this part); Part 3 reads Table 1 properly.
Table 1 as three stacks of blocks, one block per layer, block width proportional to D. ViT-Base: 12 layers of 768. ViT-Large: 24 of 1,024. ViT-Huge: 32 of 1,280. Under each, the paper's parameter count and ours without a classification head.

Where the 86 million parameters come from

Every piece of the model above has a known shape, so the count is a formula. For hidden size DD, LL layers, MLP size 4D4D, patch size PP, CC channels, NN patches and KK classes:

P2C⋅D+D⏟patch embedding E + bias+D⏟xclass+(N+1) D⏟Epos+L (4(D2+D)⏟q, k, v, Umsa+8D2+5D⏟MLP+4D⏟2 LayerNorms)+2D⏟final LN+DK+K⏟head\underbrace{P^2 C \cdot D + D}_{\text{patch embedding } \mathbf{E} \text{ + bias}} + \underbrace{D}_{\mathbf{x}_{\text{class}}} + \underbrace{(N+1)\,D}_{\mathbf{E}_{pos}} + L\,\big(\underbrace{4(D^2 + D)}_{\text{q, k, v, } U_{msa}} + \underbrace{8D^2 + 5D}_{\text{MLP}} + \underbrace{4D}_{\text{2 LayerNorms}}\big) + \underbrace{2D}_{\text{final LN}} + \underbrace{D K + K}_{\text{head}}

The per-layer part simplifies to 12D2+13D12D^2 + 13D. Each term, for ViT-B/16 at 224 with K=1,000K = 1{,}000:

  • patch embedding: 16⋅16⋅3⋅768+768=590,59216 \cdot 16 \cdot 3 \cdot 768 + 768 = 590{,}592;
  • class token: 768768; positions: 197×768=151,296197 \times 768 = 151{,}296;
  • attention per layer: 4×(7682+768)=2,362,3684 \times (768^2 + 768) = 2{,}362{,}368 (four 768×768768 \times 768 matrices with biases);
  • MLP per layer: 768⋅3,072+3,072+3,072⋅768+768=4,722,432768 \cdot 3{,}072 + 3{,}072 + 3{,}072 \cdot 768 + 768 = 4{,}722{,}432;
  • two LayerNorms per layer: 4×768=3,0724 \times 768 = 3{,}072; so one layer is 7,087,8727{,}087{,}872 and twelve are 85,054,46485{,}054{,}464;
  • final LayerNorm: 1,5361{,}536; head: 768×1,000+1,000=769,000768 \times 1{,}000 + 1{,}000 = 769{,}000.

The script evaluates the formula and counts the real model's parameters. For ViT-L/16 and ViT-H/14 it builds the model from its configuration on PyTorch's "meta" device, which creates every tensor's shape without allocating any memory, so no weights are needed:

python
def formula(D, L, P, C, N, K, mlp=None):
    mlp = mlp or 4 * D
    return (P * P * C * D + D) + D + (N + 1) * D + L * (4 * (D * D + D) + D * mlp + mlp + mlp * D + D + 4 * D) + 2 * D + (D * K + K)

c = ViTConfig(hidden_size=1280, num_hidden_layers=32, num_attention_heads=16, intermediate_size=5120, patch_size=14, num_labels=1000)
with torch.device('meta'):                       # shapes only, no memory for weights
    huge = ViTForImageClassification(c)
print(sum(p.numel() for p in huge.parameters()))
plain text
== 7. counting the parameters ==
  formula: (P^2 C D + D) + D + (N+1) D + L (12 D^2 + 13 D) + 2 D + (D K + K)
           patch emb       cls  position  L layers            LN   head
  ViT-B/16 at 224, K = 1000:
    patch embedding      590,592   [class]    768   positions   151,296
    one layer: attention 4(D^2 + D) = 2,362,368   MLP 8D^2 + 5D = 4,722,432   2 LayerNorms 4D = 3,072   -> 7,087,872   x 12 = 85,054,464
    final LayerNorm 1,536   head D K + K = 769,000
    formula total 86,567,656   counted in model.parameters() 86,567,656   same: True
    without the 1000-class head: 85,798,656   (paper, Table 1: 86M)
    where they sit: attention 28.35M (32.7%)   mlp 56.67M (65.5%)   embeddings 0.74M (0.9%)   layernorms 0.04M (0.0%)   head 0.77M (0.9%)
  ViT-L/16 at 224 (meta device, no weights): D=1024 L=24 heads=16 MLP=4096 N=196
    formula 304,326,632   counted 304,326,632   same: True
    with the 1000-class head 304.3M, without it 303.3M   (paper: 307M)
  ViT-H/14 at 224 (meta device, no weights): D=1280 L=32 heads=16 MLP=5120 N=256
    formula 632,045,800   counted 632,045,800   same: True
    with the 1000-class head 632.0M, without it 630.8M   (paper: 632M)

The formula matches the counted parameters exactly for all three sizes: 86,567,656 for ViT-B/16 with its 1,000-class head, 304,326,632 for ViT-L/16 and 632,045,800 for ViT-H/14. Compared with Table 1: ViT-Base rounds to 86M with or without the head (85.8M or 86.6M). ViT-Huge rounds to 632M with a 1,000-class head (632.0M; without it, 630.8M), so Table 1's figure for Huge appears to include a head of that size. ViT-Large is the odd one: we get 303.3M or 304.3M, and the paper says 307M. An honesty note: the 3M gap is real and we cannot explain it from the model's configuration (ViT-L/16 at 384 pixels has 577 position rows instead of 197, which adds only 0.4M; the 21,843-class ImageNet-21k head would add 22M, far too much). The released ViT-L/16 checkpoints have 304M parameters, so we trust the count and flag the table.

Where the 86,567,656 parameters of ViT-B/16 (1000-class head) sitattention (12 layers)attention (12 layers): 28.35 M (32.7%)28.35 M (32.7%)MLP (12 layers)MLP (12 layers): 56.67 M (65.5%)56.67 M (65.5%)embeddings (E, [class], E_pos)embeddings (E, [class], E_pos): 0.74 M (0.9%)0.74 M (0.9%)head (768 × 1000 + 1000)head (768 × 1000 + 1000): 0.77 M (0.9%)0.77 M (0.9%)LayerNorms (25 of them)LayerNorms (25 of them): 0.04 M (0.0%)0.04 M (0.0%)
Where the 86,567,656 parameters of ViT-B/16 sit. The MLP blocks hold 56.67 million (65.5%), attention 28.35 million (32.7%), the embeddings (E, the class token and the 197 position rows) only 0.74 million, the head 0.77 million and all 25 LayerNorms 0.04 million.

Two things stand out. The embeddings are tiny: 0.9% of the model, against about a fifth for BERT-base, because BERT carries a 30,522-row vocabulary table and ViT carries one 768×768768 \times 768 matrix. And two thirds of ViT is MLP, one third attention; the ratio 8D2:4D28D^2 : 4D^2 is fixed by the 4D4D MLP width and is the same for every Transformer with that width.

Appendix A: the attention equations

The paper keeps the attention formulas out of the main text and puts them in Appendix A. They are the inside of the "MSA" box above.

Equations 5 to 7: one head

The idea: from the same input compute queries, keys and values; score every pair; turn scores into weights; average the values with those weights.

[q,k,v]=z Uqkv,Uqkv∈RD×3Dh[\mathbf{q}, \mathbf{k}, \mathbf{v}] = \mathbf{z}\,\mathbf{U}_{qkv}, \qquad \mathbf{U}_{qkv} \in \mathbb{R}^{D \times 3D_h} A=softmax ⁣(qk⊤/Dh),A∈RN×NA = \text{softmax}\!\left(\mathbf{q}\mathbf{k}^\top / \sqrt{D_h}\right), \qquad A \in \mathbb{R}^{N \times N} SA(z)=A v\text{SA}(\mathbf{z}) = A\,\mathbf{v}

where:

  • z\mathbf{z} is the (normalised) input, N×DN \times D: 197 × 768 in ViT-B;
  • Uqkv\mathbf{U}_{qkv} is one matrix of D×3DhD \times 3D_h that produces all three at once; q,k,v\mathbf{q}, \mathbf{k}, \mathbf{v} are each N×DhN \times D_h (197 × 64). The released model stores it as three separate 768×768768 \times 768 matrices covering all 12 heads, which is the same numbers arranged differently, plus biases the equation omits;
  • qk⊤\mathbf{q}\mathbf{k}^\top is N×NN \times N: entry (i,j)(i, j) is the dot product of query ii with key jj;
  • Dh\sqrt{D_h} is 8 for Dh=64D_h = 64; dividing keeps the scores from growing with the head size, so softmax does not saturate (BERT Part 2 measures why);
  • softmax is applied to each row, so row ii holds the 197 weights of token ii and they add up to 1;
  • A vA\,\mathbf{v} is N×DhN \times D_h: row ii is the weighted average of all value rows, with token ii's weights.

This is Vaswani et al.'s Equation 1 with different letters:

A worked example small enough to check with a pencil. Four tokens, D=6D = 6, two heads, so Dh=3D_h = 3, with small whole numbers chosen so the sums stay short (vit_part2_attn.py). Head 1:

plain text
== Eq. 5 to 8 on a tiny example: N = 4 tokens, D = 6, k = 2 heads, D_h = D/k = 3 ==
  z (the input, one row per token)  shape (4, 6)
       t0      0      1      0      0      1      0
       t1      0      1      1      1      2      1
       t2      1      1      1      0      0      0
       t3      2      2      2      1      2      1
-- head 1 --
  U_qkv (head 1): D x 3 D_h  shape (6, 9)
               0      1      0      0      0      0     -1      1      0
               1     -1      1      0      1      0      0      0      1
               0      1      0     -1      1      1     -1      1      0
               1      1     -1     -1      0     -1     -1     -1      0
               1      0      1      1     -1      0      0     -1     -1
               0     -1     -1      1      0     -1      0      0     -1
  [q, k, v] = z U_qkv  (the first 3 columns are q, the next 3 k, the last 3 v)  shape (4, 9)
       t0      2     -1      2      1      0      0      0     -1      0
       t1      4      0      1      1      0     -1     -2     -2     -2
       t2      1      1      1     -1      2      1     -2      2      1
       t3      5      2      2      0      2      0     -5      1     -1
  q k^T / sqrt(D_h) = q k^T / 1.732  shape (4, 4)   (row i = query of token i, column j = key of token j)
       t0  1.155  0.000 -1.155 -1.155
       t1  2.309  1.732 -1.732  0.000
       t2  0.577  0.000  1.155  1.155
       t3  2.887  1.732  0.577  2.309
  A = softmax(row by row)  shape (4, 4)
       t0  0.661  0.208  0.066  0.066
       t1  0.596  0.335  0.010  0.059
       t2  0.195  0.110  0.348  0.348
       t3  0.506  0.160  0.050  0.284
    row sums: [1.0, 1.0, 1.0, 1.0]
  SA(z) = A v  (row i = weighted average of the value rows, weights = row i of A)  shape (4, 3)
       t0 -0.876 -0.880 -0.416
       t1 -0.986 -1.185 -0.718
       t2 -2.653  0.629 -0.219
       t3 -3.014  0.002 -0.995

Follow one number through. Row t0 of z\mathbf{z} is (0,1,0,0,1,0)(0, 1, 0, 0, 1, 0), so its query is the sum of rows 2 and 5 of Uqkv\mathbf{U}_{qkv}'s first three columns: (1,−1,1)+(1,0,1)=(2,−1,2)(1, -1, 1) + (1, 0, 1) = (2, -1, 2), as printed. Its score against its own key (1,0,0)(1, 0, 0) is 2⋅1+(−1)⋅0+2⋅0=22 \cdot 1 + (-1) \cdot 0 + 2 \cdot 0 = 2; divided by 3=1.732\sqrt{3} = 1.732 that is 1.155. Softmax over the row (1.155,0,−1.155,−1.155)(1.155, 0, -1.155, -1.155):

e1.155e1.155+e0+e−1.155+e−1.155=3.1733.173+1.000+0.315+0.315=3.1734.803=0.661\frac{e^{1.155}}{e^{1.155} + e^{0} + e^{-1.155} + e^{-1.155}} = \frac{3.173}{3.173 + 1.000 + 0.315 + 0.315} = \frac{3.173}{4.803} = 0.661

which the script checks the same way (A[t0, t0] = 3.173 / 4.803 = 0.661). And the first number of t0's output is 0.661⋅0+0.208⋅(−2)+0.066⋅(−2)+0.066⋅(−5)=−0.8760.661 \cdot 0 + 0.208 \cdot (-2) + 0.066 \cdot (-2) + 0.066 \cdot (-5) = -0.876: the weighted average of the first column of v\mathbf{v}. Token t0 listens mostly to itself (0.661) and to t1 (0.208).

Head 1 of the tiny example: scaled scores, softmax weights and SA(z) = A vq kᵀ / √3t0t1t2t3t01.150.00-1.15-1.15t12.311.73-1.730.00t20.580.001.151.15t32.891.730.582.31A = softmax (rows sum to 1)t0t1t2t3t00.660.210.070.07t10.600.330.010.06t20.200.110.350.35t30.510.160.050.28v (value rows)v₀v₁v₂t00-10t1-2-2-2t2-221t3-51-1SA(z) = A vt0-0.88-0.88-0.42t1-0.99-1.18-0.72t2-2.650.63-0.22t3-1.84-0.44-0.55row t0 of SA(z) = 0.661·v(t0) + 0.208·v(t1) + 0.066·v(t2) + 0.066·v(t3)its first number: 0.661·0 + 0.208·(-2) + 0.066·(-2) + 0.066·(-5) = -0.876Orange cells are negative numbers; the shade shows the size.
Head 1 of the tiny example as a picture. Left: the 4×4 scaled scores. Middle: softmax turns each row into weights that add up to 1. Right: the four value rows, and below them SA(z), each row a weighted average of the value rows. Orange cells are negative numbers.

Equation 8: many heads

MSA(z)=[SA1(z); SA2(z); ⋯ ; SAk(z)] Umsa,Umsa∈Rk⋅Dh×D\text{MSA}(\mathbf{z}) = [\text{SA}_1(\mathbf{z});\ \text{SA}_2(\mathbf{z});\ \cdots;\ \text{SA}_k(\mathbf{z})]\ \mathbf{U}_{msa}, \qquad \mathbf{U}_{msa} \in \mathbb{R}^{k \cdot D_h \times D}

where:

  • SAh(z)\text{SA}_h(\mathbf{z}) is head hh's output, N×DhN \times D_h, each head with its own Uqkv\mathbf{U}_{qkv};
  • the brackets place the kk outputs side by side: N×kDhN \times k D_h, which is N×DN \times D when Dh=D/kD_h = D/k;
  • Umsa\mathbf{U}_{msa} mixes the heads, kDh×Dk D_h \times D (768 × 768 in ViT-B; the released model adds a bias);
  • the result is N×DN \times D, the same shape as the input, ready for the residual addition of Equation 2.

Vaswani et al. draw it like this:

The tiny example's second head and the join:

plain text
-- Eq. 8 --
  [SA_1(z); SA_2(z)]  (the two heads side by side: 4 x 6)  shape (4, 6)
       t0 -0.876 -0.880 -0.416  0.727  0.760  0.367
       t1 -0.986 -1.185 -0.718  0.673  0.910  0.288
       t2 -2.653  0.629 -0.219  0.950  0.210  0.615
       t3 -3.014  0.002 -0.995  0.884  0.769  0.653
  U_msa: k D_h x D = 6 x 6  shape (6, 6)
               0     -1      0      0      0      0
               0      1      0     -1     -1     -1
              -1     -1      0      1      1      0
               1      0      0      0     -1     -1
               0      0      0     -1     -1     -1
               1     -1      0      1      0      1
  MSA(z) = [SA_1(z); SA_2(z)] U_msa  shape (4, 6)
       t0  1.510  0.045  0.000  0.070 -1.023 -0.240
       t1  1.679  0.231  0.000 -0.155 -1.116 -0.110
       t2  1.784  2.885  0.000 -0.442 -2.007 -1.173
       t3  2.091  1.299  0.000 -0.228 -1.766 -0.560

The first three columns of the concatenation are head 1's output from above; the last three are head 2's. Multiplying by the 6×66 \times 6 matrix Umsa\mathbf{U}_{msa} gives four rows of six numbers, the same shape as the input z\mathbf{z}. (The third column of our random Umsa\mathbf{U}_{msa} happens to be all zeros, so the third output column is zero; a trained matrix would not do that.) In ViT-B the same picture has 197 rows, 12 heads of 64, and a 768×768768 \times 768 Umsa\mathbf{U}_{msa}; the shapes figure above shows it.

Next, in Part 3: how a pre-trained ViT is adapted to a new task. Section 3.2 removes the head, attaches a zero-initialised D×KD \times K layer and fine-tunes at a higher resolution, which means more patches than the position table has rows and a 2D interpolation of the position embeddings, worked by hand. Then Section 4.1: the datasets (ImageNet, ImageNet-21k, JFT-300M), Table 1 read properly, the ResNet baselines and the hybrids, and the training recipe of Tables 3 and 4. Continue to Part 3.

Run it yourself

Every number in this part comes from two scripts in code/papers/vit/. vit_part2_math.py downloads google/vit-base-patch16-224 (about 350 MB) and google/vit-base-patch16-224-in21k, loads one picture from the huggingface/cats-image dataset, and recomputes every step of Section 3.1 and Appendix A by hand; ViT-L/16 and ViT-H/14 are built on the meta device from their configurations, so no large weights are needed. It runs on a laptop CPU in about a minute. vit_part2_attn.py is the four-token example and runs in a second.

bash
pip install torch transformers datasets pillow
python vit_part2_math.py      # prints everything below, writes results/part2_math.json
python vit_part2_attn.py      # the tiny Eq. 5 to 8 example, writes results/part2_attn.json
Terminal output of vit_part2_math.py: the picture and model configuration, the 196 patches and the first numbers of patch 0, the conv-equals-linear check, Equation 1 with the class token and position numbers, LayerNorm by hand, the attention of head 1 and head 8 of layer 1 for the class token, GELU values, the twelve layers compared with the library, Equation 4 and the top-5 classes, the head and pooler check and the parameter counts
The real output of vit_part2_math.py (first 60 lines).
Terminal output of vit_part2_attn.py: the input z, U qkv, q k v, the scores, the softmax weights and SA of z for two heads, the concatenation and U msa, the hand check of one softmax entry and the permutation equivariance test
The real output of vit_part2_attn.py (first 60 lines).

References

The ViT paper

  1. A. Dosovitskiy, L. Beyer, A. Kolesnikov, D. Weissenborn, X. Zhai, T. Unterthiner, M. Dehghani, M. Minderer, G. Heigold, S. Gelly, J. Uszkoreit, N. Houlsby. An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale (ViT). ICLR 2021.
  2. Google Research. vision_transformer: the official code and pre-trained checkpoints.
  3. Hugging Face model cards google/vit-base-patch16-224 (fine-tuned on ImageNet) and google/vit-base-patch16-224-in21k (pre-trained only), the checkpoints used in this part.

Papers the ViT paper cites in this part (the five that matter most here)

  1. A. Vaswani, N. Shazeer, N. Parmar, J. Uszkoreit, L. Jones, A. N. Gomez, Ł. Kaiser, I. Polosukhin. Attention Is All You Need (the Transformer). NeurIPS 2017.
  2. J. Devlin, M.-W. Chang, K. Lee, K. Toutanova. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. NAACL 2019.
  3. N. Parmar, A. Vaswani, J. Uszkoreit, Ł. Kaiser, N. Shazeer, A. Ku, D. Tran. Image Transformer (local self-attention). ICML 2018.
  4. J.-B. Cordonnier, A. Loukas, M. Jaggi. On the Relationship between Self-Attention and Convolutional Layers (the 2×2 patch model). ICLR 2020.
  5. A. Kolesnikov, L. Beyer, X. Zhai, J. Puigcerver, J. Yung, S. Gelly, N. Houlsby. Big Transfer (BiT): General Visual Representation Learning. ECCV 2020.

Other sources used in this part

  1. J. L. Ba, J. R. Kiros, G. E. Hinton. Layer Normalization. arXiv 2016.
  2. Code for this part: vit_part2_math.py and vit_part2_attn.py; figures by figs_part2.py, paper excerpts by shots_part2.py.