Chapter 31
Vision Transformers
Images as patches, patch embeddings, 2-D positions, and a tiny ViT.
Vision transformers treat an image as a short sequence, then reuse the same self-attention machinery that powers language models. This matters for multimodal LLMs because the visual encoder usually hands image tokens, not pixels, to a contrastive loss or a language model. The essential trick is to make patches look like tokens while preserving enough two-dimensional position information.
31.1 Images as patches
For an image batch , choose a square patch size . If divides height and width, the image becomes a grid with rows and columns. Flattening each local block gives a sequence of patch vectors:
The code is only reshape, transpose, and reshape again. The transpose moves the patch-grid axes next to each other before flattening, so tokens follow raster order: left to right, then top to bottom.
def patchify(images, patch_size):
"""Images (B, H, W, C) -> flattened patches (B, GH*GW, P*P*C)."""
b, height, width, channels = images.shape
if height % patch_size or width % patch_size:
raise ValueError("image dimensions must be divisible by patch_size")
gh, gw = height // patch_size, width // patch_size
x = images.reshape(b, gh, patch_size, gw, patch_size, channels)
x = x.transpose(0, 1, 3, 2, 4, 5)
return x.reshape(b, gh * gw, patch_size * patch_size * channels)
def sequence_length(height, width, patch_size, class_token=False):
tokens = (height // patch_size) * (width // patch_size)
return tokens + int(class_token)
Token count grows with area. A image with has patch tokens, or 197 tokens when a class token is prepended. Doubling resolution to with the same patch size gives 784 patch tokens, so self-attention is much more expensive.
Patch size is the first trade-off in a ViT. Large patches shorten the sequence and make the encoder cheap, but each token must summarize a bigger part of the image before attention can mix information. Small patches preserve fine detail, but they lengthen the sequence before any semantic reasoning has happened. That is why a patchify test is not cosmetic: one misplaced transpose silently changes which pixels a token contains.
31.2 Patch embedding and positions
A linear patch embedding maps each flattened patch into the model dimension:
This is exactly a convolution whose kernel is , whose stride is , and whose output channels are the embedding dimension. The chapter tests compare the two implementations. The convolution view explains why patch embedding is cheap: neighboring patches do not overlap, so each pixel is read once.
Thinking of patch embedding as a convolution also connects ViTs to older vision code. A framework can implement the projection with an optimized convolution kernel, then flatten the resulting feature map into a sequence. Thinking of it as a linear layer is better for book-keeping: it makes clear that a visual token is just another row in a matrix, ready for the same transformer equations as a text token.
def linear_patch_embedding(patches, weight, bias):
return patches @ weight + bias
def conv2d_strided(images, kernel, bias, stride):
"""Small NHWC convolution with kernel (P, P, C, D) and stride P."""
b, height, width, _ = images.shape
patch = kernel.shape[0]
gh, gw = height // stride, width // stride
out = np.empty((b, gh, gw, kernel.shape[-1]), dtype=images.dtype)
for r in range(gh):
for c in range(gw):
window = images[:, r * stride:r * stride + patch,
c * stride:c * stride + patch, :]
out[:, r, c, :] = np.tensordot(window, kernel,
axes=([1, 2, 3], [0, 1, 2])) + bias
return out
Transformers are permutation-equivariant unless position is added. A compact learned two-dimensional scheme stores one embedding per patch row and one per patch column, then adds their sum to the patch token at that grid location:
Absolute 2-D embeddings keep the code simple and make the grid explicit. Other models use relative or rotary variants, but every ViT needs some signal that a patch came from the top left rather than the bottom right.
The row-plus-column form is not the only learned absolute embedding, but it is easy to inspect. It says that moving down changes the row vector, moving right changes the column vector, and a location is represented by their sum. This factorization uses fewer parameters than assigning a separate vector to every grid cell and can be resized more naturally when the grid changes.
def add_2d_position(tokens, grid_hw, row_embed, col_embed):
"""Add learned row+column embeddings to patch tokens."""
gh, gw = grid_hw
pos = row_embed[:gh, None, :] + col_embed[None, :gw, :]
return tokens + pos.reshape(gh * gw, -1)[None, :, :]
def prepend_class_token(tokens, class_token):
batch = tokens.shape[0]
cls = np.broadcast_to(class_token, (batch, 1, tokens.shape[-1]))
return np.concatenate([cls, tokens], axis=1)
def pool_sequence(sequence, use_class_token=True):
return sequence[:, 0] if use_class_token else sequence.mean(axis=1)
31.3 Self-attention over visual tokens
After patch embedding, the model is just a transformer encoder. For each layer, every token builds a query, key, and value. Attention compares all query-key pairs, normalizes each row with a softmax, then mixes values:
The implementation below is deliberately self-contained: split heads, compute scaled dot products, softmax, combine heads. There is no dependency on another writer’s transformer code.
def split_heads(x, num_heads):
b, tokens, dim = x.shape
head_dim = dim // num_heads
return x.reshape(b, tokens, num_heads, head_dim).transpose(0, 2, 1, 3)
def combine_heads(x):
b, heads, tokens, head_dim = x.shape
return x.transpose(0, 2, 1, 3).reshape(b, tokens, heads * head_dim)
def multihead_self_attention(x, wq, wk, wv, wo, num_heads):
q = split_heads(x @ wq, num_heads)
k = split_heads(x @ wk, num_heads)
v = split_heads(x @ wv, num_heads)
scores = q @ k.transpose(0, 1, 3, 2) / np.sqrt(q.shape[-1])
weights = softmax(scores, axis=-1)
return combine_heads(weights @ v) @ wo, weights
The tiny forward pass uses a pre-norm residual block: layer-normalize, attend, add the residual, then layer-normalize, apply a small MLP, and add the second residual. It is enough to verify the shapes and data flow of a ViT without spending time on training.
Self-attention is the step that makes patches nonlocal. A corner patch can attend directly to a patch on the opposite corner in one layer; a convolution would need many local layers or a large kernel to connect them. The price is the score matrix. Each head forms one square matrix per image, so long visual sequences quickly dominate memory even when the embedding dimension is modest.
def transformer_block(x, params):
attn, _ = multihead_self_attention(layer_norm(x), params["wq"], params["wk"],
params["wv"], params["wo"],
params["num_heads"])
x = x + attn
hidden = gelu(layer_norm(x) @ params["mlp_w1"] + params["mlp_b1"])
return x + hidden @ params["mlp_w2"] + params["mlp_b2"]
def tiny_vit_forward(images, params, use_class_token=True):
patches = patchify(images, params["patch_size"])
tokens = linear_patch_embedding(patches, params["patch_w"], params["patch_b"])
gh = images.shape[1] // params["patch_size"]
gw = images.shape[2] // params["patch_size"]
tokens = add_2d_position(tokens, (gh, gw), params["row_pos"],
params["col_pos"])
if use_class_token:
tokens = prepend_class_token(tokens, params["class_token"])
encoded = transformer_block(tokens, params)
pooled = pool_sequence(encoded, use_class_token)
return pooled @ params["head_w"] + params["head_b"]
31.4 Class token or mean pooling
ViT classifiers need one vector for the whole image. The original pattern prepends a learned class token and reads that token after the encoder [dosovitskiy2020image]. Mean pooling uses the average of all patch tokens instead. The class token gives the model a dedicated global slot; mean pooling forces every patch representation to carry information useful to the final average. Both appear in modern vision encoders, and the better choice is empirical.
The sequence length determines memory more than the number of pixels does. With attention, each layer forms a token-by-token score matrix, so the dominant score storage scales like . Patch size is therefore a modeling decision: smaller patches preserve detail but create longer sequences.
Class-token and mean-pooling modes also behave differently under masking or cropping. A class token can learn to gather evidence from the visible tokens through attention. Mean pooling has no dedicated gatherer; every remaining patch contributes directly to the final vector. In this chapter both are just switches in the forward pass, which makes their shape consequences explicit.
|
In practice
|
ViT showed that a plain transformer encoder over fixed-size image patches can replace convolutional backbones when trained at scale [dosovitskiy2020image]. The attention block is the same scaled dot-product mechanism introduced for text transformers [vaswani2017attention]. Large vision transformers now scale to billions of parameters and are often used as frozen or lightly tuned encoders for multimodal systems [dehghani2023scaling]. Practical models spend much of their engineering budget on resolution, patch size, and token reduction because visual attention cost rises quickly with image area. |
31.5 Teach it
The one-sentence version. A vision transformer cuts an image into patches, embeds those patches as tokens, adds 2-D position information, and runs a transformer encoder.
An analogy. Treat the image like a tiled mural: each tile gets a note saying where it came from, then all tiles talk to all other tiles before the model summarizes the mural.
At the board.
-
Draw a image, mark patches, and count 196 tokens.
-
Flatten one patch and multiply by ; then show the same operation as a stride-16 convolution.
-
Add row and column position embeddings to each token.
-
Run one attention head and choose either the class token or mean pooling for the image vector.
Misconceptions to address.
-
"A ViT has no spatial bias." Patch order plus position embeddings carry spatial information.
-
"The class token is mandatory." Mean pooling is a valid alternative.
-
"Higher resolution only adds a few pixels." It can square the attention score cost.
Check for understanding. If patch size stays fixed and image height and width both double, what happens to the number of patch tokens and to the attention score matrix?
31.6 Exercises
For a image and patches, compute the patch grid, the number of patch tokens, and the sequence length with a class token. Repeat the token count for .
For a single-channel image containing values through in
row-major order, write the four flattened patches produced by patchify.
Show why a linear layer applied to flattened nonoverlapping patches is equivalent to a convolution with stride .
Run tiny_vit_forward on a synthetic batch with patches.
List the sequence length before pooling for class-token mode and mean-pooling mode, and explain
why the logits have the same shape.
References
-
[dehghani2023scaling] M. Dehghani et al. Scaling Vision Transformers to 22 Billion Parameters. 2023. arXiv:2302.05442
-
[dosovitskiy2020image] A. Dosovitskiy et al. An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. 2020. arXiv:2010.11929
-
[vaswani2017attention] A. Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762