Architectures · Vision & diffusion

Vision Transformer

ViT-B/16 · image patches as tokens

compared
ViT-B/16 vs GPT-2 small
layers
12 GPT-2 12
dmodel
768 GPT-2 768
heads
12 GPT-2 12
tokens
1 + 196 patches GPT-2 ≤ 1,024 text
output
1,000 classes GPT-2 50,257 tokens
params
86M GPT-2 124M

All steps

  1. GPT-2 vs ViT

    ViT-B/16 has GPT-2 small’s width and depth, arranged like BERT (no mask), but its tokens are image patches: the vocabulary lookup becomes one matrix product on raw pixels, and the output is one class for the whole image, read from a [CLS] token. Click a label to jump.

    12 × 768 · 86M · 197 tokens

  2. An image as a sequence

    A 224 × 224 image is cut into a 14 × 14 grid of 16 × 16 patches. Read row by row, the 196 patches become the token sequence; each one is 16 × 16 × 3 = 768 numbers of red, green and blue. (A toy grey image here.)

    224² → 14 × 14 patches → 196 tokens

  3. Patch embedding

    Each flattened patch is multiplied by one learned matrix, the same for every patch. This replaces GPT-2’s vocabulary lookup and equals a 16 × 16 convolution with stride 16. On the right, the top principal components of ViT-B/16’s real 768 filters.

    x = patches · Wpatch (768 × 768)

  4. [CLS] and positions

    A learned [CLS] token goes in front, and each of the 197 slots gets a learned position vector, a plain 1D list. Yet in the real model each patch position’s vector is most like its neighbours’ and its own row’s and column’s: training rediscovered the 2D grid. Hover a small map.

    x = [CLS; patches] + P (197 × 768)

  5. One label for the image

    All 197 tokens pass through 12 encoder layers, every patch attending to every patch. For classification only the [CLS] output is used, mapped to 1,000 ImageNet classes. ViT needed very large pretraining sets (up to 300M images) to beat convolutional networks.

    class = head(LN(xCLS))

Code

# ViT: an image becomes a sequence of patch tokens
patches = img.unfold(16)                        # (196, 16 · 16 · 3 = 768)
x = patches @ W_patch + b                       # = Conv2d(3, 768, kernel=16, stride=16)
x = torch.cat([cls_token, x]) + pos_emb         # (197, 768), learned positions
for block in blocks:                            # GPT-2 small's shape, pre-LN
    x = x + attn(ln_1(x))                       # no mask: every patch sees every patch
    x = x + mlp(ln_2(x))
logits = head(ln_f(x[0]))                       # [CLS] only → 1,000 classes

Go deeper