USAAIO Lesson 92, from Phase 3. It covers patch embedding through an nn.Conv2d with stride P, the CLS token, learned positional embeddings, applying a standard transformer encoder to the sequence of patches, and the classification head on the CLS token, then compares fine-tuning with training from scratch and the scaling behavior of ViT against a CNN. TinyViT8 is built from scratch and trained on load_digits to 96.4% test accuracy after 150 epochs, with its 26,538 parameters verified with torch 2.7.1+cpu. The lesson runs to 30 slides.
Subject: Machine Learning · 57 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 92 · Phase 3
Patches as tokens: convert an image into a sequence, pass it through a standard transformer encoder, classify on the CLS token. Every component built and verified in PyTorch — from patch embedding to attention maps.
Objectives
nn.Conv2d with stride=P and explain why it is equivalent to flattening + linear projectionWarm-up
Discussion prompt
Before we open Lesson 92: Vision Transformer (ViT): without looking back, what was the main idea of Flash Attention & Memory-Efficient Transformers, and what could you do by the end of it that you could not do before?
Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.
Answer:
Flash Attention IO complexity (O(N) HBM vs O(N²)), tiled softmax with online running-max, gradient checkpointing, ALiBi position bias, and Multi-Query Attention (MQA).
Section
Part 1 of 4
Concept
A standard transformer (Lesson 88) expects a sequence of vectors. ViT (Dosovitskiy 2021) creates that sequence by chopping the image into non-overlapping P×P patches and treating each flattened patch as a token.
\[ N = \left(\frac{H}{P}\right)^2 \quad\text{(square image } H{\times}H\text{, patch size }P\text{)} \]
| H | P | N patches | patch dim (3-ch) | seq len (+CLS) |
|---|---|---|---|---|
| 224 | 16 | 196 | 768 | 197 |
| 224 | 14 | 256 | 588 | 257 |
| 32 | 4 | 64 | 48 | 65 |
| 8 | 4 | 4 | 48 | 5 |
Comparison
Comparison matrix
From An image is worth N×P² words: refill the P column from what you know. The rest of the table is as it appeared.
| H | P | N patches | patch dim (3-ch) | seq len (+CLS) |
|---|---|---|---|---|
| 224 | 16 | 196 | 768 | 197 |
| 224 | 14 | 256 | 588 | 257 |
| 32 | 4 | 64 | 48 | 65 |
| 8 | 4 | 4 | 48 | 5 |
Concept
Each P×P×C patch is flattened to a vector of dimension P²·C, then linearly projected to the model dimension d_model. Two equivalent implementations:
nn.Linear(P²·C, d_model) — easy to understand, O(N) Python loopsnn.Conv2d(C, d_model, kernel_size=P, stride=P) — one pass over the image, all patches in parallel; output shape (B, d_model, H/P, W/P) which you then flatten and transposeThe Conv2d approach is mathematically identical to flattening + nn.Linear (same matrix multiply, just batched spatially). stride=P ensures non-overlapping patches.
Counterexample
Discussion prompt
Each P×P×C patch is flattened to a vector of dimension P²·C, then linearly projected to the model dimension d_model. Two equivalent implementations:
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
The Conv2d approach is mathematically identical to flattening + nn.Linear (same matrix multiply, just batched spatially). stride=P ensures non-overlapping patches.
Estimation
Predict first
Trace the patch embed layer on a batch of 4 images of size (1, 8, 8) with P=4, d_model=32. Predict the output shape at each step before reading on.
Commit before you compute: what does Patch embedding forward pass — tensor shapes come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Conv2d output: (4, 32, 2, 2) — 2×2 spatial grid of feature maps, one per patch
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. With kernel=4, stride=4 on an 8×8 input: out_size = (8-4)/4+1 = 2; there are 2×2=4 non-overlapping patch positions.
Worked example
Trace the patch embed layer on a batch of 4 images of size (1, 8, 8) with P=4, d_model=32. Predict the output shape at each step before reading on.
import torch, torch.nn as nn
class PatchEmbedding(nn.Module):
def __init__(self, img_size, patch_size, in_chans, d_model):
super().__init__()
self.proj = nn.Conv2d(
in_chans, d_model,
kernel_size=patch_size, stride=patch_size
)
def forward(self, x):
x = self.proj(x) # (B, d_model, H/P, W/P)
x = x.flatten(2) # (B, d_model, N)
x = x.transpose(1, 2) # (B, N, d_model)
return x
torch.manual_seed(42)
patch_embed = PatchEmbedding(img_size=8, patch_size=4, in_chans=1, d_model=32)
x = torch.randn(4, 1, 8, 8)
out = patch_embed(x)
print(out.shape) # torch.Size([4, 4, 32])Conv2d output: (4, 32, 2, 2) — 2×2 spatial grid of feature maps, one per patch
Why: With kernel=4, stride=4 on an 8×8 input: out_size = (8-4)/4+1 = 2; there are 2×2=4 non-overlapping patch positions.
| operation | tensor shape | meaning |
|---|---|---|
| input x | (4, 1, 8, 8) | B=4 images, C=1, H=W=8 |
| after Conv2d | (4, 32, 2, 2) | B=4, d_model=32, 2×2 patch grid |
| after flatten(2) | (4, 32, 4) | merge spatial dims → 4 patches |
| after transpose | (4, 4, 32) | B=4, N=4 patches, d=32 ← final |
Trade off
Comparison matrix
From Patch embedding forward pass — tensor shapes: every row here is a choice with a cost. Fill the tensor shape column, then say which row you would actually pick and what you give up for it.
| operation | tensor shape | meaning |
|---|---|---|
| input x | (4, 1, 8, 8) | B=4 images, C=1, H=W=8 |
| after Conv2d | (4, 32, 2, 2) | B=4, d_model=32, 2×2 patch grid |
| after flatten(2) | (4, 32, 4) | merge spatial dims → 4 patches |
| after transpose | (4, 4, 32) | B=4, N=4 patches, d=32 ← final |
Section
Part 2 of 4
Concept
A ViT prepends a learnable [CLS] token to the patch sequence before the transformer encoder. It has no corresponding image patch — its role is to accumulate global information through attention.
After the encoder, only the CLS token's output vector is passed to the classification head. This is the same design as BERT (Lesson 90): the sequence is [CLS, patch_1, patch_2, ..., patch_N].
\[ \mathbf{z}_0 = [\mathbf{x}_{\text{cls}};\; E\mathbf{p}_1;\; E\mathbf{p}_2;\; \ldots;\; E\mathbf{p}_N] + \mathbf{E}_{\text{pos}} \]
Analogy
Discussion prompt
Explain The CLS token by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
A ViT prepends a learnable [CLS] token to the patch sequence before the transformer encoder. It has no corresponding image patch — its role is to accumulate global information through attention.
Concept
Unlike the sinusoidal scheme, ViT uses learned positional embeddings — a parameter matrix of shape (N+1, d_model) added to the token sequence. They are trained end-to-end and encode 2D position implicitly.
(1, N+1, d_model) — broadcast over the batchtrunc_normal_(std=0.02)), learned during training| embedding type | shape (N=4, d=32) | parameters | 2D awareness |
|---|---|---|---|
| learned pos embed | (1, 5, 32) | 160 | implicit (learned) |
| sinusoidal (Lesson 88) | (1, 5, 32) | 0 (fixed) | 1D only by default |
| 2D sin/cos extension | (1, 5, 32) | 0 (fixed) | explicit row/col freqs |
Estimation
Predict first
Continue from the patch embed output (4, 4, 32). Add a CLS token and positional embedding. Trace shapes at each step.
Commit before you compute: what does CLS prepend + positional embedding come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: cls_token.expand(B, -1, -1): shape (4, 1, 32) — replicate across batch, keep d and seq-len dims
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The single learned CLS vector is shared across the batch; expand avoids allocating extra memory (no copy, just a view).
Worked example
Continue from the patch embed output (4, 4, 32). Add a CLS token and positional embedding. Trace shapes at each step.
import torch, torch.nn as nn
B, N, d = 4, 4, 32
patch_out = torch.randn(B, N, d) # simulated patch embed output
cls_token = nn.Parameter(torch.zeros(1, 1, d))
pos_embed = nn.Parameter(torch.zeros(1, N + 1, d)) # N+1 for CLS
nn.init.trunc_normal_(cls_token, std=0.02)
nn.init.trunc_normal_(pos_embed, std=0.02)
cls = cls_token.expand(B, -1, -1) # (4, 1, 32)
x = torch.cat([cls, patch_out], dim=1) # (4, 5, 32)
x = x + pos_embed # broadcast add
print(x.shape) # torch.Size([4, 5, 32])cls_token.expand(B, -1, -1): shape (4, 1, 32) — replicate across batch, keep d and seq-len dims
Why: The single learned CLS vector is shared across the batch; expand avoids allocating extra memory (no copy, just a view).
| tensor | shape | note |
|---|---|---|
| patch_out | (4, 4, 32) | N=4 patch tokens, d=32 |
| cls expanded | (4, 1, 32) | same CLS for all batch items |
| after cat dim=1 | (4, 5, 32) | N+1=5 tokens total |
| after +pos_embed | (4, 5, 32) | adds (1,5,32) broadcast |
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at CLS prepend + positional embedding. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Anomaly
Predict first
A student writes this, and it looks reasonable:
The CLS token is appended after the patch tokens: [patch_1, ..., patch_N, CLS]. The last position's output goes to the head.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: ViT prepends the CLS token — it is always at index 0 in the sequence.
ViT prepends CLS: [CLS, patch_1, ..., patch_N]. The zeroth output of the encoder is the classification token.
Why: ViT prepends the CLS token — it is always at index 0 in the sequence. Appending to the end is a design choice used in some models, but original ViT always uses position 0.
Trap
The CLS token is appended after the patch tokens: [patch_1, ..., patch_N, CLS]. The last position's output goes to the head.
Use x[:, -1] as the classification vector
Why: This is wrong. ViT prepends the CLS token — it is always at index 0 in the sequence. Appending to the end is a design choice used in some models, but original ViT always uses position 0.
ViT prepends CLS: [CLS, patch_1, ..., patch_N]. The zeroth output of the encoder is the classification token.
Use x[:, 0] as the classification vector
Why: In code: cls_out = encoder_output[:, 0]. Because CLS was prepended at dim=1 index 0, the first element of the encoder output sequence is always the CLS representation.
Break the constraint
Discussion prompt
The rule this trap just fixed:
In code: cls_out = encoder_output[:, 0]. Because CLS was prepended at dim=1 index 0, the first element of the encoder output sequence is always the CLS representation.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
ViT prepends the CLS token — it is always at index 0 in the sequence. Appending to the end is a design choice used in some models, but original ViT always uses position 0.
Section
Part 3 of 4
Concept
ViT reuses the standard transformer encoder block unchanged from Lesson 88 — no modifications needed. Each block: LayerNorm → MultiHeadAttention → residual → LayerNorm → MLP → residual.
4 × d_modelExplain it
Discussion prompt
Explain Standard transformer encoder (recap Lesson 88) to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.
Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.
Answer:
ViT reuses the standard transformer encoder block unchanged from Lesson 88 — no modifications needed. Each block: LayerNorm → MultiHeadAttention → residual → LayerNorm → MLP → residual.
Estimation
Predict first
Build TinyViT8 (8×8 images, P=4, d_model=32, 2 heads, 2 encoder blocks, 10 classes). This is a runnable ViT — no pretrained weights, no downloads.
Commit before you compute: what does TinyViT8 — full forward pass with shapes come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: 26,538 parameters; output shape (4, 10) — batch of 4, 10 logits
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Patch projection: Conv2d(1,32,4,4)=544 params.
Worked example
Build TinyViT8 (8×8 images, P=4, d_model=32, 2 heads, 2 encoder blocks, 10 classes). This is a runnable ViT — no pretrained weights, no downloads.
import torch, torch.nn as nn
class TransformerBlock(nn.Module):
def __init__(self, d, n_heads):
super().__init__()
self.norm1 = nn.LayerNorm(d)
self.attn = nn.MultiheadAttention(d, n_heads, batch_first=True)
self.norm2 = nn.LayerNorm(d)
mlp_d = d * 4
self.mlp = nn.Sequential(
nn.Linear(d, mlp_d), nn.GELU(), nn.Linear(mlp_d, d)
)
def forward(self, x):
a, _ = self.attn(self.norm1(x), self.norm1(x), self.norm1(x))
x = x + a
x = x + self.mlp(self.norm2(x))
return x
class TinyViT8(nn.Module):
def __init__(self):
super().__init__()
d, P, n_heads = 32, 4, 2
self.patch_proj = nn.Conv2d(1, d, kernel_size=P, stride=P)
n_p = (8 // P) ** 2 # 4 patches
self.cls_token = nn.Parameter(torch.zeros(1, 1, d))
self.pos_embed = nn.Parameter(torch.zeros(1, n_p + 1, d))
self.blocks = nn.ModuleList([TransformerBlock(d, n_heads) for _ in range(2)])
self.norm = nn.LayerNorm(d)
self.head = nn.Linear(d, 10)
def forward(self, x):
B = x.shape[0]
x = self.patch_proj(x).flatten(2).transpose(1, 2) # (B,4,32)
x = torch.cat([self.cls_token.expand(B,-1,-1), x], 1) + self.pos_embed
for blk in self.blocks:
x = blk(x)
return self.head(self.norm(x)[:, 0])
torch.manual_seed(42)
m = TinyViT8()
print(sum(p.numel() for p in m.parameters()), 'params')
print(m(torch.randn(4, 1, 8, 8)).shape)26,538 parameters; output shape (4, 10) — batch of 4, 10 logits
Why: Patch projection: Conv2d(1,32,4,4)=544 params. CLS=32, pos_embed=160. Each TransformerBlock contributes ~12,832 params (attention + MLP). Two blocks = 25,664. LayerNorm + head = 138. Total verified: 26,538.
| component | parameter count | notes |
|---|---|---|
| patch_proj (Conv2d) | 544 | 1×32×4×4 weights + 32 bias |
| cls_token | 32 | learnable, broadcast over batch |
| pos_embed | 160 | (1, 5, 32): 4 patches + CLS |
| block 0 (attn+MLP) | 12,832 | attn 4,224 + MLP 8,576 + LN 128 |
| block 1 (attn+MLP) | 12,832 | same structure |
| norm + head | 138 | LN 64 + Linear(32,10)=330... 74 |
| TOTAL | 26,538 | verified with named_parameters() |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
26,538 parameters; output shape (4, 10) — batch of 4, 10 logits
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Build TinyViT8 (8×8 images, P=4, d_model=32, 2 heads, 2 encoder blocks, 10 classes). This is a runnable ViT — no pretrained weights, no downloads.
Anomaly
Predict first
A student writes this, and it looks reasonable:
After the encoder, pass the entire sequence through the head: logits = head(encoder_out.mean(dim=1)).
It is wrong. Say what breaks — and say it before you turn the page.
Correct: This is wrong for ViT-style architectures.
Extract only the CLS token (index 0) and pass it to the head: cls_out = encoder_out[:, 0].
Why: This is wrong for ViT-style architectures. Mean-pooling all tokens does work (it is called 'GAP ViT'), but original ViT uses only the CLS token representation, and mixing them wastes the specialized CLS mechanism.
Trap
After the encoder, pass the entire sequence through the head: logits = head(encoder_out.mean(dim=1)).
Average all token outputs and classify — treating ViT like a mean-pool backbone
Why: This is wrong for ViT-style architectures. Mean-pooling all tokens does work (it is called 'GAP ViT'), but original ViT uses only the CLS token representation, and mixing them wastes the specialized CLS mechanism.
Extract only the CLS token (index 0) and pass it to the head: cls_out = encoder_out[:, 0].
logits = head(norm(encoder_out)[:, 0])
Why: The CLS token was prepended at position 0 to gather global image information through attention. The patch tokens at positions 1..N are discarded at classification time.
Two truths and a lie
Sort into buckets
Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.
P×P×C patch is flattened to a vector of dimension P²·C, then linearly projected to the model dimension d_model. Two equivalent implementations:; Build rules: type every line; print tensor shapes after each operation; verify the CLS index is 0 before classification; use torch.manual_seed(42) for reproducibility.[patch_1, ..., patch_N, CLS]. The last position's output goes to the head.; After the encoder, pass the entire sequence through the head: logits = head(encoder_out.mean(dim=1)).Section
Part 4 of 4
Concept
CNNs have strong inductive biases baked in: translation equivariance (Lesson 55) and locality. A CNN achieves good accuracy on small datasets because the architecture already knows 'nearby pixels matter more.'
ViT has no such inductive bias — it treats all patch pairs equally and must learn spatial structure from scratch. This means ViT needs large datasets (JFT-300M, ImageNet-21k) to match CNN performance.
| regime | < 10k images | 100k–1M images | 100M+ images |
|---|---|---|---|
| ResNet/CNN | competitive | competitive | strong but saturates |
| ViT-Base | weaker | approaches CNN | matches/beats CNN |
| ViT-Large/H | weaker | approaches CNN | outperforms CNN |
Comparison
Comparison matrix
From ViT needs more data — but scales better: refill the 100M+ images column from what you know. The rest of the table is as it appeared.
| regime | < 10k images | 100k–1M images | 100M+ images |
|---|---|---|---|
| ResNet/CNN | competitive | competitive | strong but saturates |
| ViT-Base | weaker | approaches CNN | matches/beats CNN |
| ViT-Large/H | weaker | approaches CNN | outperforms CNN |
Concept
In practice, ViT models are pretrained on massive datasets (ImageNet-21k or JFT) and then fine-tuned on the target task — the same transfer-learning workflow as BERT (Lesson 90).
Analogy
Discussion prompt
Explain Fine-tuning vs training from scratch by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
In practice, ViT models are pretrained on massive datasets (ImageNet-21k or JFT) and then fine-tuned on the target task — the same transfer-learning workflow as BERT (Lesson 90).
Constraint
Discussion prompt
Run The ViT architecture recipe with this step confiscated:
Positional embed: add learned (1, N+1, d_model) parameter — broadcast over batch
Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.
Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.
Answer:
N = (H/P)²; sequence length = N+1 (CLS prepended). For 224×224, P=16 → N=196.nn.Conv2d(C, d_model, kernel_size=P, stride=P) → flatten(2).transpose(1,2) → shape (B, N, d_model)torch.cat([cls_token.expand(B,-1,-1), patches], dim=1) → (B, N+1, d_model)(1, N+1, d_model) parameter — broadcast over batch[PreNorm → MHA → residual → PreNorm → MLP(4d) → residual]; no causal masknorm(x)[:, 0] → Linear(d_model, n_classes) — index 0 = CLS token outputPattern
N = (H/P)²; sequence length = N+1 (CLS prepended). For 224×224, P=16 → N=196.nn.Conv2d(C, d_model, kernel_size=P, stride=P) → flatten(2).transpose(1,2) → shape (B, N, d_model)torch.cat([cls_token.expand(B,-1,-1), patches], dim=1) → (B, N+1, d_model)(1, N+1, d_model) parameter — broadcast over batch[PreNorm → MHA → residual → PreNorm → MLP(4d) → residual]; no causal masknorm(x)[:, 0] → Linear(d_model, n_classes) — index 0 = CLS token outputEdge cases
Discussion prompt
The ViT architecture recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
N = (H/P)²; sequence length = N+1 (CLS prepended). For 224×224, P=16 → N=196.nn.Conv2d(C, d_model, kernel_size=P, stride=P) → flatten(2).transpose(1,2) → shape (B, N, d_model)torch.cat([cls_token.expand(B,-1,-1), patches], dim=1) → (B, N+1, d_model)(1, N+1, d_model) parameter — broadcast over batch[PreNorm → MHA → residual → PreNorm → MLP(4d) → residual]; no causal masknorm(x)[:, 0] → Linear(d_model, n_classes) — index 0 = CLS token outputElimination
Eliminate the wrong options
A ViT processes 224×224 images with patch size P=16. What is the total sequence length fed into the transformer encoder (including the CLS token)?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: N = (224/16)² = 14² = 196 patches. Prepend 1 CLS token → sequence length = 197. This is the standard ViT-Base/16 setting verified in the patch count table.
Check
Work it out before clicking.
Check your understanding
A ViT processes 224×224 images with patch size P=16. What is the total sequence length fed into the transformer encoder (including the CLS token)?
Answer: A
Why: N = (224/16)² = 14² = 196 patches. Prepend 1 CLS token → sequence length = 197. This is the standard ViT-Base/16 setting verified in the patch count table.
Prediction
Predict first
Why is nn.Conv2d(C, d_model, kernel_size=P, stride=P) equivalent to flattening each P×P×C patch and applying nn.Linear(P²·C, d_model)?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: Both compute the same linear map: each kernel of shape (d_model, C, P, P) multiplies one non-overlapping patch, producing one d_model-dimensional output vector per patch.
Why: A Conv2d with kernel_size=P and stride=P slides a P×P×C filter exactly to each non-overlapping patch position with no overlap. Each application is a dot product of the flattened patch with a weight row — identical to nn.Linear(P²·C, d_model) applied to the flattened patch. The two have the same number of parameters: d_model × (C × P²) weights + d_model biases.
Check
Think about what Conv2d with stride=P computes.
Check your understanding
Why is nn.Conv2d(C, d_model, kernel_size=P, stride=P) equivalent to flattening each P×P×C patch and applying nn.Linear(P²·C, d_model)?
Answer: A
Why: A Conv2d with kernel_size=P and stride=P slides a P×P×C filter exactly to each non-overlapping patch position with no overlap. Each application is a dot product of the flattened patch with a weight row — identical to nn.Linear(P²·C, d_model) applied to the flattened patch. The two have the same number of parameters: d_model × (C × P²) weights + d_model biases.
Elimination
Eliminate the wrong options
You have a dataset of 5,000 labeled medical images. Which model is likely to achieve higher accuracy with standard fine-tuning from ImageNet pretrained weights?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: At 5,000 images, CNNs' inductive biases (locality, translation equivariance) give them a decisive advantage over ViT. ViT-Large pretraining on 21k can help, but ResNet-50 fine-tuned from ImageNet typically outperforms or matches ViT at small target-dataset sizes. Training any ViT from scratch on 5k images will almost certainly underfit badly.
Check
Apply the inductive bias reasoning.
Check your understanding
You have a dataset of 5,000 labeled medical images. Which model is likely to achieve higher accuracy with standard fine-tuning from ImageNet pretrained weights?
Answer: A
Why: At 5,000 images, CNNs' inductive biases (locality, translation equivariance) give them a decisive advantage over ViT. ViT-Large pretraining on 21k can help, but ResNet-50 fine-tuned from ImageNet typically outperforms or matches ViT at small target-dataset sizes. Training any ViT from scratch on 5k images will almost certainly underfit badly.
Section
Project
Concept
Build TinyViT8 — a minimal but structurally correct ViT — on load_digits (1,797 samples, 8×8 grayscale, 10 classes). Three milestones: patch shape tracing → full forward pass → training loop.
| # | milestone | key tool |
|---|---|---|
| 1 | Implement PatchEmbedding, trace shapes through Conv2d→flatten→transpose | nn.Conv2d, tensor.flatten, tensor.transpose |
| 2 | Add CLS token + pos embed, verify sequence shape (B, N+1, d) | nn.Parameter, torch.cat |
| 3 | Train 150 epochs, target >95% test accuracy | Adam lr=3e-4, CrossEntropyLoss |
Build rules: type every line; print tensor shapes after each operation; verify the CLS index is 0 before classification; use torch.manual_seed(42) for reproducibility.
Counterexample
Discussion prompt
Build TinyViT8 — a minimal but structurally correct ViT — on load_digits (1,797 samples, 8×8 grayscale, 10 classes). Three milestones: patch shape tracing → full forward pass → training loop.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
Build rules: type every line; print tensor shapes after each operation; verify the CLS index is 0 before classification; use torch.manual_seed(42) for reproducibility.
Worked example
Your turn: implement PatchEmbedding for 8×8 images with P=4, d=32. Predict the output shape before running. What is out_size = (8 - 4) // 4 + 1?
Hint: Conv2d with kernel_size=P, stride=P produces (B, d_model, H/P, W/P). Then flatten(2) merges spatial dims, transpose(1,2) swaps channels and sequence.
import torch, torch.nn as nn
class PatchEmbedding(nn.Module):
def __init__(self, img_size, patch_size, in_chans, d_model):
super().__init__()
self.proj = nn.Conv2d(
in_chans, d_model, kernel_size=patch_size, stride=patch_size
)
def forward(self, x):
x = self.proj(x) # (B, d_model, H/P, W/P)
x = x.flatten(2) # (B, d_model, N)
return x.transpose(1, 2) # (B, N, d_model)
torch.manual_seed(42)
pe = PatchEmbedding(8, 4, 1, 32)
x = torch.randn(4, 1, 8, 8)
out = pe(x)
print(out.shape) # torch.Size([4, 4, 32])| step | shape | formula |
|---|---|---|
| input | (4, 1, 8, 8) | B=4, C=1, H=W=8 |
| after Conv2d | (4, 32, 2, 2) | (8-4)/4+1 = 2 per axis |
| after flatten(2) | (4, 32, 4) | 2×2 = 4 patches merged |
| after transpose | (4, 4, 32) | N=4, d=32 ← final output |
Worked example
Your turn: prepend the CLS token and add positional embeddings. Predict the final sequence shape before running. Why must pos_embed have shape (1, N+1, d_model) and not (1, N, d_model)?
Hint: cls_token.expand(B, -1, -1) broadcasts the (1,1,d) parameter to (B,1,d) without copying. torch.cat([cls, patches], dim=1) then gives (B, N+1, d).
import torch, torch.nn as nn
B, N, d = 4, 4, 32
patches = torch.randn(B, N, d)
cls_token = nn.Parameter(torch.zeros(1, 1, d))
pos_embed = nn.Parameter(torch.zeros(1, N + 1, d))
nn.init.trunc_normal_(cls_token, std=0.02)
nn.init.trunc_normal_(pos_embed, std=0.02)
cls = cls_token.expand(B, -1, -1) # (4, 1, 32)
seq = torch.cat([cls, patches], dim=1) # (4, 5, 32)
seq = seq + pos_embed # (4, 5, 32)
print(seq.shape, ' N+1 =', N+1)| tensor | shape | why |
|---|---|---|
| cls expanded | (4, 1, 32) | broadcast single CLS across batch |
| after cat(dim=1) | (4, 5, 32) | N=4 patches + 1 CLS = 5 |
| pos_embed | (1, 5, 32) | must cover all 5 positions |
| seq + pos_embed | (4, 5, 32) | (1,5,32) broadcasts to (4,5,32) |
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at Milestone 2 — CLS + positional embedding. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Worked example
Your turn: complete the training loop for 150 epochs with Adam lr=3e-4. Predict how quickly ViT converges compared to the CNN from Lesson 55.
Hint: use the same five-step loop as Lesson 40 (zero_grad → forward → loss → backward → step). A lower learning rate (3e-4 vs 1e-3) is typical for transformers — they are sensitive to lr.
import torch, torch.nn as nn
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
import numpy as np
digits = load_digits()
X = digits.data.reshape(-1,1,8,8).astype(np.float32)/16.0
y = digits.target.astype(np.int64)
Xtr,Xte,ytr,yte = train_test_split(X,y,test_size=0.2,random_state=42)
Xt=torch.tensor(Xtr); yt=torch.tensor(ytr)
Xtt=torch.tensor(Xte); ytt=torch.tensor(yte)
torch.manual_seed(42)
model = TinyViT8() # defined in earlier milestone
opt = torch.optim.Adam(model.parameters(), lr=3e-4)
loss_fn = nn.CrossEntropyLoss()
for epoch in range(150):
opt.zero_grad()
loss = loss_fn(model(Xt), yt)
loss.backward(); opt.step()
with torch.no_grad():
acc=(model(Xtt).argmax(1)==ytt).float().mean()
print(f'test acc: {acc.item()*100:.1f}%') # 96.4%| epoch | train loss | test acc |
|---|---|---|
| 0 | 2.4068 | 11.9% |
| 9 | 2.1581 | 21.9% |
| 24 | 1.6183 | 50.6% |
| 49 | 1.0804 | 77.5% |
| 99 | 0.5458 | 91.1% |
| 149 | 0.2860 | 96.4% |
Trade off
Comparison matrix
From Milestone 3 — train to >95% accuracy: every row here is a choice with a cost. Fill the train loss column, then say which row you would actually pick and what you give up for it.
| epoch | train loss | test acc |
|---|---|---|
| 0 | 2.4068 | 11.9% |
| 9 | 2.1581 | 21.9% |
| 24 | 1.6183 | 50.6% |
| 49 | 1.0804 | 77.5% |
| 99 | 0.5458 | 91.1% |
| 149 | 0.2860 | 96.4% |
Concept
import torch, torch.nn as nn
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
import numpy as np
class PatchEmbedding(nn.Module):
def __init__(self, img_size, patch_size, in_chans, d):
super().__init__()
self.proj = nn.Conv2d(in_chans, d, kernel_size=patch_size, stride=patch_size)
def forward(self, x):
return self.proj(x).flatten(2).transpose(1, 2)
class Block(nn.Module):
def __init__(self, d, h):
super().__init__()
self.n1 = nn.LayerNorm(d); self.n2 = nn.LayerNorm(d)
self.attn = nn.MultiheadAttention(d, h, batch_first=True)
self.mlp = nn.Sequential(nn.Linear(d, d*4), nn.GELU(), nn.Linear(d*4, d))
def forward(self, x):
a,_ = self.attn(self.n1(x), self.n1(x), self.n1(x))
x = x + a; x = x + self.mlp(self.n2(x)); return x
class TinyViT8(nn.Module):
def __init__(self, d=32, n_heads=2, n_layers=2, n_cls=10):
super().__init__()
self.pe = PatchEmbedding(8, 4, 1, d)
self.cls = nn.Parameter(torch.zeros(1,1,d))
self.pos = nn.Parameter(torch.zeros(1,5,d)) # 4 patches + CLS
self.blocks = nn.ModuleList([Block(d, n_heads) for _ in range(n_layers)])
self.norm = nn.LayerNorm(d); self.head = nn.Linear(d, n_cls)
nn.init.trunc_normal_(self.cls, std=0.02); nn.init.trunc_normal_(self.pos, std=0.02)
def forward(self, x):
B=x.shape[0]; x=self.pe(x)
x=torch.cat([self.cls.expand(B,-1,-1),x],1)+self.pos
for b in self.blocks: x=b(x)
return self.head(self.norm(x)[:,0])
digits=load_digits(); X=digits.data.reshape(-1,1,8,8).astype(np.float32)/16.0; y=digits.target.astype(np.int64)
Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=0.2,random_state=42)
Xt=torch.tensor(Xtr); yt=torch.tensor(ytr)
torch.manual_seed(42); m=TinyViT8()
opt=torch.optim.Adam(m.parameters(),lr=3e-4); L=nn.CrossEntropyLoss()
for _ in range(150):
opt.zero_grad(); L(m(Xt),yt).backward(); opt.step()
with torch.no_grad():
acc=(m(torch.tensor(Xte)).argmax(1)==torch.tensor(yte)).float().mean()
print(f'{acc.item()*100:.1f}%') # 96.4%| design choice | value | rationale |
|---|---|---|
| patch size P | 4 (on 8×8) | gives N=4 patches — minimal but valid |
| d_model | 32 | small enough to train in seconds on CPU |
| n_heads | 2 | d_k = 32/2 = 16 per head |
| learning rate | 3e-4 | transformers need lower lr than CNNs |
| epochs | 150 | slower to converge than CNN (Lesson 55) |
| test accuracy | 96.4% | vs 97.2% for 2-conv CNN at 100 epochs |
Comparison
Comparison matrix
From The full program: refill the value column from what you know. The rest of the table is as it appeared.
| design choice | value | rationale |
|---|---|---|
| patch size P | 4 (on 8×8) | gives N=4 patches — minimal but valid |
| d_model | 32 | small enough to train in seconds on CPU |
| n_heads | 2 | d_k = 32/2 = 16 per head |
| learning rate | 3e-4 | transformers need lower lr than CNNs |
| epochs | 150 | slower to converge than CNN (Lesson 55) |
| test accuracy | 96.4% | vs 97.2% for 2-conv CNN at 100 epochs |
Concept
Out loud, slides closed: (1) explain the full ViT forward pass from raw image pixels to logits, naming every tensor shape at each stage; (2) explain why the CLS token is at index 0 and not the end; (3) state the inductive bias argument for why ViT underperforms CNN on small datasets.
Stretch (homework from the lesson plan): visualize attention maps — after training, extract attn_weights from nn.MultiheadAttention using need_weights=True, reshape to (N+1, N+1), and inspect which patches the CLS token attends to. Compare ViT vs the SmallCNN from Lesson 55 at 1k vs 10k training examples.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Patches as tokens — the core idea · CLS token & positional embeddings · Transformer encoder on patch sequences · ViT vs CNN — when to use which · Your turn: TinyViT8 from scratch. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
N = (H/P)² patches and sequence length N+1 for any ViT configurationnn.Conv2d(C, d, k=P, s=P) → flatten(2).transpose(1,2)nn.MultiheadAttention + pre-norm MLP blocks| component | the one thing to remember |
|---|---|
| patch count | N = (H/P)²; seq len = N+1 after CLS prepend |
| patch embed | Conv2d(k=P, s=P) is identical to flatten + Linear |
| CLS token | prepended at index 0; only its output goes to the head |
| pos embed | learned (N+1, d_model) parameter; added before encoder |
| ViT vs CNN | ViT < CNN at < ~10k images; ViT ≥ CNN at ~100M+ samples |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.