feat(phase-07/09): vision transformers (ViT)

This commit is contained in:
Rohit Ghumare
2026-04-23 00:12:04 +01:00
parent 7008f4b8dc
commit 043f335cfd
5 changed files with 419 additions and 0 deletions
@@ -0,0 +1,102 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 900 540" font-family="Georgia, 'Times New Roman', serif">
<defs>
<marker id="arrow" viewBox="0 0 10 10" refX="9" refY="5" markerWidth="6" markerHeight="6" orient="auto">
<path d="M0,0 L10,5 L0,10 z" fill="#1a1a1a"/>
</marker>
<style>
.box { fill: #faf6ef; stroke: #1a1a1a; stroke-width: 1.5; }
.hot { fill: #fff1d6; stroke: #c0392b; stroke-width: 1.5; }
.cls { fill: #2c3e50; }
.cls-text { fill: #fff; }
.patch { fill: #e7edd9; stroke: #1a1a1a; stroke-width: 0.6; }
.label { font-size: 14px; font-weight: 600; fill: #1a1a1a; }
.content { font-size: 12px; fill: #333; font-family: 'Menlo', monospace; }
.caption { font-size: 11px; fill: #555; font-style: italic; }
.title { font-size: 16px; font-weight: 700; fill: #1a1a1a; }
</style>
</defs>
<text x="450" y="30" text-anchor="middle" class="title">an image is a grid of patches — nothing else changes</text>
<!-- Original image 4x4 grid -->
<text x="130" y="70" text-anchor="middle" class="label">image (224x224x3)</text>
<g transform="translate(60 85)">
<g class="patch">
<rect x="0" y="0" width="36" height="36"/>
<rect x="36" y="0" width="36" height="36"/>
<rect x="72" y="0" width="36" height="36"/>
<rect x="108" y="0" width="36" height="36"/>
<rect x="0" y="36" width="36" height="36"/>
<rect x="36" y="36" width="36" height="36"/>
<rect x="72" y="36" width="36" height="36"/>
<rect x="108" y="36" width="36" height="36"/>
<rect x="0" y="72" width="36" height="36"/>
<rect x="36" y="72" width="36" height="36"/>
<rect x="72" y="72" width="36" height="36"/>
<rect x="108" y="72" width="36" height="36"/>
<rect x="0" y="108" width="36" height="36"/>
<rect x="36" y="108" width="36" height="36"/>
<rect x="72" y="108" width="36" height="36"/>
<rect x="108" y="108" width="36" height="36"/>
</g>
</g>
<text x="130" y="250" text-anchor="middle" class="caption">sliced into a 14 × 14 grid of 16-pixel patches</text>
<line x1="230" y1="160" x2="270" y2="160" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
<!-- Flatten into tokens -->
<text x="400" y="70" text-anchor="middle" class="label">flatten each patch → embed</text>
<g transform="translate(280 90)">
<rect x="0" y="0" width="26" height="40" class="cls"/><text x="13" y="25" text-anchor="middle" class="content cls-text">CLS</text>
<rect x="30" y="0" width="26" height="40" class="patch"/><text x="43" y="25" text-anchor="middle" class="content">p_1</text>
<rect x="60" y="0" width="26" height="40" class="patch"/><text x="73" y="25" text-anchor="middle" class="content">p_2</text>
<rect x="90" y="0" width="26" height="40" class="patch"/><text x="103" y="25" text-anchor="middle" class="content">p_3</text>
<rect x="120" y="0" width="26" height="40" class="patch"/><text x="133" y="25" text-anchor="middle" class="content">p_4</text>
<text x="160" y="25" class="content">...</text>
<rect x="180" y="0" width="36" height="40" class="patch"/><text x="198" y="25" text-anchor="middle" class="content">p_196</text>
</g>
<text x="400" y="155" text-anchor="middle" class="caption">sequence of 197 d-dim tokens (CLS + 196 patches)</text>
<line x1="400" y1="170" x2="400" y2="200" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
<!-- Transformer encoder block -->
<rect x="240" y="200" width="320" height="190" class="box"/>
<text x="400" y="225" text-anchor="middle" class="label">transformer encoder (same as BERT)</text>
<rect x="270" y="245" width="260" height="28" class="box"/>
<text x="400" y="264" text-anchor="middle" class="content">block 1: self-attn + MLP</text>
<rect x="270" y="280" width="260" height="28" class="box"/>
<text x="400" y="299" text-anchor="middle" class="content">block 2: self-attn + MLP</text>
<text x="400" y="325" text-anchor="middle" class="caption">× 12 for ViT-Base, × 24 for ViT-Large</text>
<rect x="270" y="340" width="260" height="32" class="hot"/>
<text x="400" y="361" text-anchor="middle" class="content">197 × d_model hidden states</text>
<line x1="400" y1="390" x2="400" y2="415" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
<!-- Classification head -->
<rect x="275" y="415" width="250" height="30" class="box"/>
<text x="400" y="435" text-anchor="middle" class="content">take CLS → Linear(d, num_classes)</text>
<line x1="400" y1="445" x2="400" y2="465" stroke="#1a1a1a" stroke-width="1.5" marker-end="url(#arrow)"/>
<rect x="305" y="465" width="190" height="30" class="hot"/>
<text x="400" y="485" text-anchor="middle" class="content">softmax → class</text>
<!-- Right side: variant callouts -->
<rect x="620" y="90" width="240" height="120" class="box"/>
<text x="740" y="112" text-anchor="middle" class="label">variants</text>
<text x="740" y="134" text-anchor="middle" class="caption">DeiT (2021): distill from CNN</text>
<text x="740" y="152" text-anchor="middle" class="caption">Swin: shifted-window local attn</text>
<text x="740" y="170" text-anchor="middle" class="caption">DINOv2: self-supervised</text>
<text x="740" y="188" text-anchor="middle" class="caption">SigLIP: + text encoder</text>
<text x="740" y="206" text-anchor="middle" class="caption">SAM 3: + mask decoder</text>
<rect x="620" y="230" width="240" height="160" class="box"/>
<text x="740" y="252" text-anchor="middle" class="label">scale table</text>
<text x="740" y="274" text-anchor="middle" class="content">ViT-B/16 86M</text>
<text x="740" y="294" text-anchor="middle" class="content">ViT-L/16 304M</text>
<text x="740" y="314" text-anchor="middle" class="content">ViT-H/14 632M</text>
<text x="740" y="334" text-anchor="middle" class="content">ViT-22B 22,000M</text>
<text x="740" y="362" text-anchor="middle" class="caption">data hunger: ≥ImageNet-21k</text>
<text x="740" y="378" text-anchor="middle" class="caption">needed to beat ResNet</text>
<text x="450" y="528" text-anchor="middle" class="caption">one architecture, many modalities — patchify the input and attention does the rest.</text>
</svg>

After

Width:  |  Height:  |  Size: 6.1 KiB

@@ -0,0 +1,147 @@
"""Vision Transformer (ViT) — the patchify + embed front end.
Pure stdlib. Takes a toy 24x24x3 image, cuts it into 6x6 patches,
projects each to a d_model vector, prepends [CLS], adds 2D position.
Verifies shapes and counts parameters for real ViT configs.
"""
import math
import random
def make_image(H, W, C=3, seed=0):
rng = random.Random(seed)
return [[[rng.randint(0, 255) / 255.0 for _ in range(C)] for _ in range(W)] for _ in range(H)]
def patchify(image, patch_size):
H = len(image)
W = len(image[0])
C = len(image[0][0])
assert H % patch_size == 0 and W % patch_size == 0
patches = []
grid = []
for row_idx, i in enumerate(range(0, H, patch_size)):
grid_row = []
for col_idx, j in enumerate(range(0, W, patch_size)):
patch = []
for di in range(patch_size):
for dj in range(patch_size):
patch.extend(image[i + di][j + dj])
patches.append(patch)
grid_row.append((row_idx, col_idx))
grid.append(grid_row)
return patches, (H // patch_size, W // patch_size)
def linear_project(patches, d_model, rng=None):
if rng is None:
rng = random.Random(0)
in_dim = len(patches[0])
scale = math.sqrt(2.0 / (in_dim + d_model))
W = [[rng.gauss(0, scale) for _ in range(d_model)] for _ in range(in_dim)]
out = []
for patch in patches:
row = [0.0] * d_model
for i, x in enumerate(patch):
if x == 0.0:
continue
for j in range(d_model):
row[j] += x * W[i][j]
out.append(row)
return out, W
def cls_and_pos(tokens, grid_h, grid_w, rng=None):
"""Prepend learnable [CLS] and add 2D sinusoidal positional encoding."""
if rng is None:
rng = random.Random(1)
d_model = len(tokens[0])
cls = [rng.gauss(0, 0.02) for _ in range(d_model)]
pe = pos_2d(grid_h, grid_w, d_model)
out = [list(cls)]
idx = 0
for i in range(grid_h):
for j in range(grid_w):
t = [tokens[idx][k] + pe[i][j][k] for k in range(d_model)]
out.append(t)
idx += 1
return out
def pos_2d(H, W, d_model):
"""2D sinusoidal: split d_model in half, encode row and col independently."""
assert d_model % 4 == 0, "d_model must be divisible by 4 for 2D sinusoidal"
half = d_model // 2
pe = [[[0.0] * d_model for _ in range(W)] for _ in range(H)]
for i in range(H):
for j in range(W):
for k in range(half // 2):
theta_row = i / (10000 ** (2 * k / half))
pe[i][j][2 * k] = math.sin(theta_row)
pe[i][j][2 * k + 1] = math.cos(theta_row)
for k in range(half // 2):
theta_col = j / (10000 ** (2 * k / half))
pe[i][j][half + 2 * k] = math.sin(theta_col)
pe[i][j][half + 2 * k + 1] = math.cos(theta_col)
return pe
def param_count_vit(d_model, n_layers, n_heads, ffn_expansion, num_patches, num_classes):
"""Approximate ViT parameter count (patch embed + transformer + head)."""
# Patch embedding: (patch_flat_size, d_model) — ignore patch_size here, caller scales.
# Self-attention per layer: 4 * d_model^2 (Q,K,V,O)
# FFN per layer: 2 * d_model * (ffn_expansion * d_model)
# Norms: 2 * d_model per layer (LayerNorm gamma+beta)
per_layer = 4 * d_model ** 2 + 2 * d_model * int(ffn_expansion * d_model) + 4 * d_model
# Position embeddings: (num_patches + 1) * d_model
pos_emb = (num_patches + 1) * d_model
# CLS token: d_model
# Classifier head: d_model * num_classes
head = d_model * num_classes
# Final layer norm: 2 * d_model
return per_layer * n_layers + pos_emb + d_model + head + 2 * d_model
def main():
H, W, C = 24, 24, 3
patch_size = 6
d_model = 48
image = make_image(H, W, C, seed=0)
patches, grid = patchify(image, patch_size)
tokens, W_proj = linear_project(patches, d_model, rng=random.Random(42))
tokens_with_pos = cls_and_pos(tokens, grid[0], grid[1], rng=random.Random(7))
print("=== ViT front-end sanity ===")
print(f"image: ({H}, {W}, {C})")
print(f"patch size: {patch_size}x{patch_size}")
print(f"grid: {grid[0]} x {grid[1]} = {grid[0] * grid[1]} patches")
print(f"flat patch size: {patch_size * patch_size * C}")
print(f"d_model: {d_model}")
print(f"sequence length: {len(tokens_with_pos)} (patches + CLS)")
print(f"cell [0,0] of CLS: {tokens_with_pos[0][0]:.4f}")
print(f"cell [0,0] of p1: {tokens_with_pos[1][0]:.4f}")
print()
print("=== parameter counts (approximate) ===")
for name, d, L, H_heads, exp, patch in [
("ViT-Tiny/16", 192, 12, 3, 4, 16),
("ViT-Small/16", 384, 12, 6, 4, 16),
("ViT-Base/16", 768, 12, 12, 4, 16),
("ViT-Large/16", 1024, 24, 16, 4, 16),
("ViT-Huge/14", 1280, 32, 16, 4, 14),
]:
grid_n = (224 // patch) ** 2
params = param_count_vit(d, L, H_heads, exp, grid_n, num_classes=1000)
# Add patch embed: (P*P*3) * d_model
params += patch * patch * 3 * d
print(f" {name:<14} d={d:<5} L={L:<3} heads={H_heads:<3} patches={grid_n:<4} ~{params / 1e6:.1f}M params")
print()
print("takeaway: vit reuses the bert encoder verbatim; all the vision smarts live")
print("in patchify + positional scheme + [cls] pooling.")
if __name__ == "__main__":
main()
@@ -0,0 +1,152 @@
# Vision Transformers (ViT)
> An image is a grid of patches. A sentence is a grid of tokens. The same transformer eats both.
**Type:** Build
**Languages:** Python
**Prerequisites:** Phase 7 · 05 (Full Transformer), Phase 4 · 03 (CNNs), Phase 4 · 14 (Vision Transformers intro)
**Time:** ~45 minutes
## The Problem
Before 2020, computer vision meant convolutions. Every SOTA on ImageNet, COCO, and detection benchmarks used a CNN backbone. Transformers were for language.
Dosovitskiy et al. (2020) — "An Image is Worth 16x16 Words" — showed you can drop the convolutions entirely. Slice an image into fixed-size patches, linearly project each patch into an embedding, feed the sequence to a vanilla transformer encoder. At sufficient scale (ImageNet-21k pretraining or bigger), ViT matches or beats ResNet-based models.
ViT was the start of a broader pattern in 2026: one architecture, many modalities. Whisper tokenizes audio. ViT tokenizes images. Action tokens for robotics. Pixel tokens for video. The transformer doesn't care — feed it a sequence and it learns.
By 2026, ViT and its descendants (DeiT, Swin, DINOv2, ViT-22B, SAM 3) own most of vision. CNNs still win on edge devices and latency-sensitive tasks. Everything else has a ViT somewhere in the stack.
## The Concept
![Image → patches → tokens → transformer](../assets/vit.svg)
### Step 1 — patchify
Split a `H × W × C` image into an `N × (P·P·C)` sequence of flat patches. Typical setup: `224 × 224` image, `16 × 16` patches → 196 patches of 768 values each.
```
image (224, 224, 3) → 14 × 14 grid of 16x16x3 patches → 196 vectors of length 768
```
Patch size is the lever. Smaller patches = more tokens, better resolution, quadratic attention cost. Larger patches = coarser, cheaper.
### Step 2 — linear embedding
A single learned matrix projects each flat patch to `d_model`. Equivalent to a convolution of kernel size `P` and stride `P`. In PyTorch this is literally `nn.Conv2d(C, d_model, kernel_size=P, stride=P)` — a 2-line implementation.
### Step 3 — prepend `[CLS]` token, add positional embeddings
- Prepend a learnable `[CLS]` token. Its final hidden state is the image representation used for classification.
- Add learnable positional embeddings (ViT-original) or sinusoidal 2D (later variants).
- In 2024+ RoPE extended to 2D for position, sometimes without explicit embeddings.
### Step 4 — standard transformer encoder
Stack L blocks of `LayerNorm → Self-Attention → + → LayerNorm → MLP → +`. Identical to BERT. No vision-specific layers. This is the pedagogical punchline of the paper.
### Step 5 — head
For classification: take `[CLS]` hidden state → linear → softmax. For DINOv2 or SAM, discard `[CLS]`, use the patch embeddings directly.
### Variants that mattered
| Model | Year | Change |
|-------|------|--------|
| ViT | 2020 | The original. Fixed patch size, full global attention. |
| DeiT | 2021 | Distillation; trainable on ImageNet-1k only. |
| Swin | 2021 | Hierarchical with shifted windows. Fixed sub-quadratic cost. |
| DINOv2 | 2023 | Self-supervised (no labels). Best general vision features. |
| ViT-22B | 2023 | 22B params; scaling laws apply. |
| SigLIP | 2023 | ViT + language pair, sigmoid contrastive loss. |
| SAM 3 | 2025 | Segment anything; ViT-Large + promptable mask decoder. |
### Why it took a while
ViT needs *a lot* of data to match CNNs because it has none of the CNN inductive biases (translation invariance, locality). Without >100M labeled images or strong self-supervised pretraining, CNNs still win at matched compute. DeiT fixed this in 2021 with distillation tricks; DINOv2 fixed it permanently in 2023 with self-supervision.
## Build It
See `code/main.py`. Pure-stdlib patchify + linear embedding + sanity checks. No training — ViT at any realistic scale needs PyTorch and hours of GPU time.
### Step 1: fake image
A 24 × 24 RGB image as a list of rows of `(R, G, B)` tuples. We use 6×6 patches → 16 patches, 108-d embedding vector each.
### Step 2: patchify
```python
def patchify(image, P):
H = len(image)
W = len(image[0])
patches = []
for i in range(0, H, P):
for j in range(0, W, P):
patch = []
for di in range(P):
for dj in range(P):
patch.extend(image[i + di][j + dj])
patches.append(patch)
return patches
```
Raster order: row-major across the grid. Every ViT uses this ordering.
### Step 3: linear embed
Multiply each flat patch by a random `(patch_flat_size, d_model)` matrix. Verify output shape is `(N_patches + 1, d_model)` after prepending `[CLS]`.
### Step 4: count parameters for a realistic ViT
Print the param count for ViT-Base: 12 layers, 12 heads, d=768, patch=16. Compare to ResNet-50 (~25M). ViT-Base lands at ~86M. ViT-Large ~307M. ViT-Huge ~632M.
## Use It
```python
from transformers import ViTImageProcessor, ViTModel
import torch
from PIL import Image
processor = ViTImageProcessor.from_pretrained("google/vit-base-patch16-224-in21k")
model = ViTModel.from_pretrained("google/vit-base-patch16-224-in21k")
img = Image.open("cat.jpg")
inputs = processor(img, return_tensors="pt")
out = model(**inputs).last_hidden_state # (1, 197, 768): [CLS] + 196 patches
cls_emb = out[:, 0] # image representation
```
**DINOv2 embeddings are the 2026 default for image features.** Freeze the backbone, train a tiny head. Works for classification, retrieval, detection, captioning. Meta's DINOv2 checkpoints outperform CLIP on every non-text vision task.
**Patch-size picking.** Small models use 16×16 (ViT-B/16). Dense prediction (segmentation) uses 8×8 or 14×14 (SAM, DINOv2). Very large models use 14×14.
## Ship It
See `outputs/skill-vit-configurator.md`. The skill picks a ViT variant and patch size for a new vision task given dataset size, resolution, and compute budget.
## Exercises
1. **Easy.** Run `code/main.py`. Verify the number of patches equals `(H/P) * (W/P)` and the flat patch dimension equals `P*P*C`.
2. **Medium.** Implement 2D sinusoidal positional embeddings — two independent sinusoidal codes for `row` and `col` of each patch, concatenated. Feed them into a tiny PyTorch ViT and compare accuracy vs learnable positional embeddings on CIFAR-10.
3. **Hard.** Build a 3-layer ViT (PyTorch), train on 1,000 MNIST images with 4×4 patches. Measure test accuracy. Now add DINOv2 pretraining on the same 1,000 images (simplified: just train the encoder to predict patch embeddings from masked patches). Does accuracy improve?
## Key Terms
| Term | What people say | What it actually means |
|------|-----------------|-----------------------|
| Patch | "The vision-transformer token" | Flat vector of pixel values for a `P × P × C` region of the image. |
| Patchify | "Chop + flatten" | Slice image into non-overlapping patches, flatten each to a vector. |
| `[CLS]` token | "The image summary" | Prepended learnable token; its final embedding is the image representation. |
| Inductive bias | "What the model assumes" | ViT has fewer priors than CNNs; needs more data to make up the gap. |
| DINOv2 | "Self-supervised ViT" | Trained without labels using image augmentation + momentum teacher. Best general image features in 2026. |
| SigLIP | "CLIP's successor" | ViT + text encoder trained with sigmoid contrastive loss; better than CLIP on matched compute. |
| Swin | "Windowed ViT" | Hierarchical ViT with local attention + shifted windows; sub-quadratic. |
| Register tokens | "2023 trick" | A few extra learnable tokens that soak up attention sinks; improves DINOv2 features. |
## Further Reading
- [Dosovitskiy et al. (2020). An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale](https://arxiv.org/abs/2010.11929) — the ViT paper.
- [Touvron et al. (2021). Training data-efficient image transformers & distillation through attention](https://arxiv.org/abs/2012.12877) — DeiT.
- [Liu et al. (2021). Swin Transformer: Hierarchical Vision Transformer using Shifted Windows](https://arxiv.org/abs/2103.14030) — Swin.
- [Oquab et al. (2023). DINOv2: Learning Robust Visual Features without Supervision](https://arxiv.org/abs/2304.07193) — DINOv2.
- [Darcet et al. (2023). Vision Transformers Need Registers](https://arxiv.org/abs/2309.16588) — the register-token fix for DINOv2.
@@ -0,0 +1,18 @@
---
name: vit-configurator
description: Pick a ViT variant, patch size, and pretraining source for a new vision task.
version: 1.0.0
phase: 7
lesson: 9
tags: [transformers, vit, vision]
---
Given a vision task (classification / segmentation / detection / retrieval), image resolution, dataset size (labeled + unlabeled), and deployment target, output:
1. Backbone. One of: DINOv2 ViT-L/14 (default for retrieval/classification), SAM 3 encoder (segmentation), SigLIP (vision-language), ConvNeXt (latency-critical). One-sentence reason.
2. Patch size. 16 for standard classification at 224, 14 for DINOv2, 8 for dense prediction at high res. Flag sequence length `(H/P)^2 + 1` and attention cost `O(N^2)`.
3. Pretraining source. Checkpoint name. For small labeled sets (<10k): DINOv2 features frozen + linear probe. For >100k: fine-tune last blocks. State why.
4. Training recipe. Optimizer (AdamW), lr, augmentations (RandAug, MixUp, Random Erasing), label smoothing (0.1 typical), EMA.
5. Risk note. Data regime risk (too little data for full fine-tune), resolution mismatch (pretrain 224 → deploy 1024 without position interpolation), register-token absence (may hurt DINOv2 features).
Refuse to recommend training a ViT from scratch on less than 1M images — CNN baselines will win. Refuse to recommend patch size that yields sequence length > 4096 without explicit discussion of Flash Attention + hierarchical variants (Swin). Flag any deployment that changes input resolution without interpolating positional embeddings.