Architectures · Vision & diffusion
- 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
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 tokensAn 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 tokensPatch 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)[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)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