mirror of
https://github.com/rohitg00/ai-engineering-from-scratch.git
synced 2026-10-02 01:54:39 +08:00
feat(phase-07/09): vision transformers (ViT)
This commit is contained in:
@@ -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
|
||||
|
||||

|
||||
|
||||
### 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.
|
||||
+18
@@ -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.
|
||||
Reference in New Issue
Block a user