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:
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.
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. 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, L 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.
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 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.
where H=W=224 is the resolution the model was trained at, P=16 is the patch size (the "/16" in "ViT-B/16"), and C=3 for red, green and blue. So the picture becomes xp, a matrix of 196 rows and 768 columns. A happy coincidence of ViT-B/16 at 224: the row length P2C=768 equals D=768, so E happens to be square. For ViT-B/32 the rows are 32⋅32⋅3=3,072 long and E is 3,072×768.
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:
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] and then mapped to [−1,1] by (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 n sits at grid row ⌊n/14⌋ and column nmod14, 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 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. Each row of xp is multiplied by the same matrix E∈R(P2⋅C)×D. One matrix product does all 196 patches at once:
196×768xp⋅768×768E=196×768xpEThe 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×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, flattenedours = patches @ E + conv.bias # Eq. 1's x_p E, for all 196 patchestheirs = conv(pixel_values).flatten(2).transpose(1, 2)[0] # what the library computesprint((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−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.
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 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.
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.
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.
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.
Check the second row by hand: patch 1's embedding started 0.058,−0.024,−0.248,3.689,… (step 2 above); add Epos[1]=0.156,−0.124,0.447,0.009,… and you get 0.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,… in both), so z0[0] is close to 2xclass. 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.
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.
where μ is the mean of the 768 numbers, σ their standard deviation (with a tiny ϵ=10−12 added so that it is never zero), and γ,β 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 numbersln = model.vit.layers[0].layernorm_beforemu, 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, then 0.065×4.321+0.007=0.288, against the printed 0.287 (the inputs shown are rounded). Notice how small γ is here (0.03 to 0.19): the trained model scales the normalised numbers down a lot before attention sees them.
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×768 matrix is multiplied by three matrices to give queries, keys and values, each 197×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 and turned into weights by softmax, and the token's output is the weighted sum of the 197 value rows. The 12 outputs (197×64 each) are placed side by side into 197×768 and multiplied by one more 768×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.attentionq = ln1 @ at.q_proj.weight.T + at.q_proj.bias # (197, 768); the same for k and vqh, 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 1SA = A @ vh # (12, 197, 64)concat = SA.transpose(0, 1).reshape(197, 768) # the 12 heads side by sidemsa = concat @ at.o_proj.weight.T + at.o_proj.bias # U_msaz_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 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 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]
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.
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
where:
zℓ′ is the output of Equation 2, 197×768;
LN is a second LayerNorm with its own γ,β;
MLP(x)=GELU(xW1+b1)W2+b2 with W1 of 768×3,072 and W2 of 3,072×768, applied to each of the 197 rows separately;
zℓ is the output of layer ℓ, 197×768, the input of the next layer.
And GELU itself, number by number:
GELU(x)=x⋅Φ(x),Φ(x)=21(1+erf(x/2))
where Φ(x) is the probability that a standard normal random number is below x: close to 0 for very negative x, 0.5 at 0, close to 1 for large x. 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:
GELU(−1)=−1×0.1587=−0.159, not 0 as ReLU would give; GELU(2)=2×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:
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). Our z1 matches the library's output of layer 1 to 6.7×10−6, and our attention weights match its weights to 2.7×10−6. So Equations 2 and 3, as written above, are the complete layer.
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−6 in layer 1 to 1.8×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 is about 1,700, against 21 in z1), 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−6 in every layer, and the final class scores agree to 3×10−6.
The idea in one sentence: after the last layer, normalise the class token one more time; that vector is the picture.
y=LN(zL0)
where zL0 is row 0 (the class token) of the last layer's output, 768 numbers, and LN is a final LayerNorm with its own γ,β. y, also 768 numbers, is the image representation. The 196 patch rows of zL are computed and then ignored. The head is not in the equation; in the fine-tuned model it is logits=yWhead+bhead with Whead of 768×1000:
python
y = model.vit.layernorm(zL[0]) # Eq. 4, 768 numberslogits = 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.
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,W2 for every token; the attention is "global" because every row of the 197×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.
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.
In shapes: a ResNet turns the 224×224×3 picture into a 14×14×1,024 feature map (with the extended-stage-3 option the paper pairs with ViT-B/16). Reading each position as a 1×1 "patch" gives N=196 tokens of 1,024 numbers, so E becomes 1,024×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.
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.
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 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.
Every piece of the model above has a known shape, so the count is a formula. For hidden size D, L layers, MLP size 4D, patch size P, C channels, N patches and K classes:
patch embedding E + biasP2C⋅D+D+xclassD+Epos(N+1)D+L(q, k, v, Umsa4(D2+D)+MLP8D2+5D+2 LayerNorms4D)+final LN2D+headDK+K
The per-layer part simplifies to 12D2+13D. Each term, for ViT-B/16 at 224 with K=1,000:
patch embedding: 16⋅16⋅3⋅768+768=590,592;
class token: 768; positions: 197×768=151,296;
attention per layer: 4×(7682+768)=2,362,368 (four 768×768 matrices with biases);
MLP per layer: 768⋅3,072+3,072+3,072⋅768+768=4,722,432;
two LayerNorms per layer: 4×768=3,072; so one layer is 7,087,872 and twelve are 85,054,464;
final LayerNorm: 1,536; head: 768×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 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×768 matrix. And two thirds of ViT is MLP, one third attention; the ratio 8D2:4D2 is fixed by the 4D MLP width and is the same for every Transformer with that width.
z is the (normalised) input, N×D: 197 × 768 in ViT-B;
Uqkv is one matrix of D×3Dh that produces all three at once; q,k,v are each N×Dh (197 × 64). The released model stores it as three separate 768×768 matrices covering all 12 heads, which is the same numbers arranged differently, plus biases the equation omits;
qk⊤ is N×N: entry (i,j) is the dot product of query i with key j;
Dh is 8 for Dh=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 i holds the 197 weights of token i and they add up to 1;
Av is N×Dh: row i is the weighted average of all value rows, with token i'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=6, two heads, so Dh=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 is (0,1,0,0,1,0), so its query is the sum of rows 2 and 5 of Uqkv's first three columns: (1,−1,1)+(1,0,1)=(2,−1,2), as printed. Its score against its own key (1,0,0) is 2⋅1+(−1)⋅0+2⋅0=2; divided by 3=1.732 that is 1.155. Softmax over the row (1.155,0,−1.155,−1.155):
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.876: the weighted average of the first column of v. Token t0 listens mostly to itself (0.661) and to t1 (0.208).
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.
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×6 matrix Umsa gives four rows of six numbers, the same shape as the input z. (The third column of our random Umsa 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×768Umsa; 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×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 pillowpython vit_part2_math.py # prints everything below, writes results/part2_math.jsonpython vit_part2_attn.py # the tiny Eq. 5 to 8 example, writes results/part2_attn.json
The real output of vit_part2_math.py (first 60 lines).The real output of vit_part2_attn.py (first 60 lines).
Papers the ViT paper cites in this part (the five that matter most here)
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.