USAAIO Lesson 63, from Week 22 of Phase 3. It covers the fine-tuning strategies, running from a linear probe through gradual unfreezing to a full fine-tune, then domain adaptation and distribution shift, few-shot learning with 5-shot tasks and prototypical networks, and zero-shot learning through textual descriptions, which is the foundation of CLIP. All the numbers were verified on the digits dataset with torch 2.7.1 and sklearn in June 2026. The lesson runs to 31 slides.
Subject: Machine Learning · 58 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 63 · Week 22
Reuse what a pretrained model already knows: linear probe → gradual unfreeze → full fine-tune. Then few-shot with prototypical networks and zero-shot with CLIP-style text embeddings.
Objectives
Warm-up
Discussion prompt
Before we open Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods: without looking back, what was the main idea of Learning Rate Schedules, 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:
linear warmup, cosine annealing, SGDR cosine-with-restarts, the one-cycle policy (LR + momentum), and the inverse-square-root Transformer schedule — each derived from first principles, verified in PyTorch LambdaLR / built-in schedulers, with a full LambdaLR from-scratch implementation project.
Section
Part 1 of 4
Concept
A network trained on a large source domain (ImageNet, large text corpora) encodes general structure — edges, textures, syntax — in its backbone layers. The final head (classifier) is task-specific and shallow.
| layer type | what it encodes | reusable? |
|---|---|---|
| early backbone | low-level features (edges, curves) | almost always |
| deep backbone | semantic concepts (dog, wheel) | usually, with care |
| head (linear) | class boundaries for source task | no — replace it |
Transfer learning = freeze the backbone's knowledge, replace the head for the new task. Later, you may unfreeze the backbone with a smaller learning rate.
Comparison
Comparison matrix
From The pretrained backbone as a feature extractor: refill the what it encodes column from what you know. The rest of the table is as it appeared.
| layer type | what it encodes | reusable? |
|---|---|---|
| early backbone | low-level features (edges, curves) | almost always |
| deep backbone | semantic concepts (dog, wheel) | usually, with care |
| head (linear) | class boundaries for source task | no — replace it |
Concept
Distribution shift means the source domain (ImageNet photos) and target domain (chest X-rays) follow different feature distributions. The backbone may transfer well (edges still exist) but the head is useless.
| source → target | shift severity | strategy |
|---|---|---|
| ImageNet → pet photos | low | linear probe usually sufficient |
| ImageNet → medical images | medium | gradual unfreeze + domain data |
| natural images → satellite | high | full fine-tune with larger dataset |
| vision → language | extreme | do NOT transfer; train from scratch |
Trade off
Comparison matrix
From Domain adaptation: source → target: every row here is a choice with a cost. Fill the shift severity column, then say which row you would actually pick and what you give up for it.
| source → target | shift severity | strategy |
|---|---|---|
| ImageNet → pet photos | low | linear probe usually sufficient |
| ImageNet → medical images | medium | gradual unfreeze + domain data |
| natural images → satellite | high | full fine-tune with larger dataset |
| vision → language | extreme | do NOT transfer; train from scratch |
Intuition
If you have only 5-50 labeled examples per class, fine-tuning all parameters overfits catastrophically — there is not enough signal to update millions of weights.
A linear probe (freeze the backbone entirely, train only the new head) has far fewer parameters (d × C where d = embedding dim, C = classes) and is safe to fit on tiny labeled sets. The pretrained backbone does the heavy lifting.
Once the head has oriented correctly, you can gradually unfreeze deeper layers with a smaller learning rate — letting the backbone adapt without catastrophic forgetting.
Counterexample
Discussion prompt
If you have only 5-50 labeled examples per class, fine-tuning all parameters overfits catastrophically — there is not enough signal to update millions of weights.
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:
Once the head has oriented correctly, you can gradually unfreeze deeper layers with a smaller learning rate — letting the backbone adapt without catastrophic forgetting.
Worked example
Setup: pretrain a 2-layer net on digits (64 → 32 ReLU → 10) for 300 epochs on all 1437 training samples (backbone frozen for transfer). Then sample 5 examples per class (50 total) and compare strategies.
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
digits = load_digits()
X = digits.data.astype('float32') / 16.0
y = digits.target
Xtr, Xte, ytr, yte = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y)
class TinyNet(nn.Module):
def __init__(self):
super().__init__()
self.backbone = nn.Sequential(nn.Linear(64, 32), nn.ReLU())
self.head = nn.Linear(32, 10)
def forward(self, x): return self.head(self.backbone(x))
def features(self, x): return self.backbone(x)
torch.manual_seed(0)
pretrained = TinyNet()
opt = torch.optim.Adam(pretrained.parameters(), lr=0.01)
for _ in range(300):
opt.zero_grad()
F.cross_entropy(pretrained(torch.tensor(Xtr)), torch.tensor(ytr)).backward()
opt.step()| epoch | training loss |
|---|---|
| 0 | 2.3273 |
| 50 | 0.1411 |
| 100 | 0.0575 |
| 200 | 0.0185 |
| 300 | 0.0084 |
Pattern
Step through it
Step through Linear probe vs. from-scratch on 5-shot one row at a time. What is driving the change, and what would the row after the last one be?
Worked example
Sample 5 examples per class (50 total). Measure test accuracy of (A) training from scratch, (B) linear probe (frozen backbone), (C) full fine-tune from pretrained checkpoint.
def few_shot(X, y, k=5, seed=7):
rng = np.random.default_rng(seed)
idx = []
for c in np.unique(y):
ci = np.where(y == c)[0]
idx.extend(rng.choice(ci, k, replace=False).tolist())
return X[idx], y[idx]
Xk, yk = few_shot(Xtr, ytr, k=5) # 50 samples
Xk_t = torch.tensor(Xk); yk_t = torch.tensor(yk)
# A: from scratch
torch.manual_seed(1)
scratch = TinyNet()
opt_s = torch.optim.Adam(scratch.parameters(), lr=0.01)
for _ in range(200):
opt_s.zero_grad()
F.cross_entropy(scratch(Xk_t), yk_t).backward(); opt_s.step()
# B: linear probe (freeze backbone)
torch.manual_seed(1)
probe = TinyNet(); probe.load_state_dict(pretrained.state_dict())
for p in probe.backbone.parameters(): p.requires_grad = False
opt_p = torch.optim.Adam(probe.head.parameters(), lr=0.05)
for _ in range(200):
opt_p.zero_grad()
F.cross_entropy(probe(Xk_t), yk_t).backward(); opt_p.step()| strategy | test acc |
|---|---|
| from scratch (5-shot) | 0.8611 |
| linear probe (5-shot) | 0.9556 |
| full fine-tune (5-shot) | 0.9583 |
| pretrained (all data) | 0.9667 |
Anomaly
Predict first
A student writes this, and it looks reasonable:
We have a pretrained backbone and 50 labeled samples. Jump straight to Adam(model.parameters(), lr=0.01) and train for 300 epochs.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: The randomly initialized head produces large gradients in the first few steps.
Stage 1: freeze the backbone, train only the head at a larger lr=0.05 for 100 epochs.
Why: The randomly initialized head produces large gradients in the first few steps. These large gradients propagate through the entire backbone and overwrite the pretrained representations — catastrophic forgetting.
Trap
We have a pretrained backbone and 50 labeled samples. Jump straight to Adam(model.parameters(), lr=0.01) and train for 300 epochs.
Unfreeze everything immediately at a large learning rate
Why: The randomly initialized head produces large gradients in the first few steps. These large gradients propagate through the entire backbone and overwrite the pretrained representations — catastrophic forgetting.
Stage 1: freeze the backbone, train only the head at a larger lr=0.05 for 100 epochs.
Stage 1 — linear probe at lr=0.05 for 100 epochs
Why: Head gradients are now bounded by a sensible loss before any backbone weight gets touched. Stage 1 acc = 0.9583 (digits 5-shot).
Stage 2 — unfreeze backbone at lr=0.001 for 100 more epochs
Why: The backbone fine-tunes gently from a stable initialization. The lower lr prevents overwriting the pretrained representations.
Break the constraint
Discussion prompt
The rule this trap just fixed:
Head gradients are now bounded by a sensible loss before any backbone weight gets touched. Stage 1 acc = 0.9583 (digits 5-shot).
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:
The randomly initialized head produces large gradients in the first few steps. These large gradients propagate through the entire backbone and overwrite the pretrained representations — catastrophic forgetting.
Section
Part 2 of 4
Concept
Few-shot learning is the formal setting where at inference time you see only k labeled examples per novel class (the support set) and must classify a query set of unlabeled samples.
| term | meaning | example |
|---|---|---|
| N-way | number of classes | 10-way = 10 classes |
| k-shot | examples per class in support | 5-shot = 5 per class |
| support set | labeled examples provided at test | 5-shot × 10 = 50 samples |
| query set | unlabeled samples to classify | the actual test queries |
| episode | one (support, query) pair for training | meta-learning iterate |
Concept
Prototypical networks (Snell et al. 2017) map all support examples through a learned embedding function f_φ, then represent each class by its mean embedding (the prototype).
\[ c_k = \frac{1}{|S_k|} \sum_{(x_i,y_i)\in S_k} f_\phi(x_i) \]
A query x is classified by computing its distance to each prototype and applying a softmax.
\[ p(y=k \mid x) = \frac{\exp(-d(f_\phi(x),\, c_k))}{\sum_{k'} \exp(-d(f_\phi(x),\, c_{k'}))} \]
Analogy
Discussion prompt
Explain Prototypical networks: class prototypes 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:
Prototypical networks (Snell et al. 2017) map all support examples through a learned embedding function f_φ, then represent each class by its mean embedding (the prototype).
Estimation
Predict first
Setup: 3 classes, 3-shot (9 support vectors in 4-dim embedding space). Compute each prototype and classify a query point at the origin.
Commit before you compute: what does Prototype computation — explicit trace come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Predicted class = 1 (nearest prototype, dist = 0.6611)
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 query at the origin is closest to class 1's prototype in L2 distance — the decision rule is argmin over distances, equivalent to argmax of the softmax above.
Worked example
Setup: 3 classes, 3-shot (9 support vectors in 4-dim embedding space). Compute each prototype and classify a query point at the origin.
import torch
torch.manual_seed(42)
# 3 classes x 3 shots, 4-dim embedding
support = torch.randn(9, 4)
labels = torch.tensor([0,0,0, 1,1,1, 2,2,2])
# compute class prototypes
protos = torch.stack([
support[labels == c].mean(dim=0) for c in range(3)
])
query = torch.zeros(1, 4) # query at origin
dists = torch.cdist(query, protos) # Euclidean distances
pred = dists.argmin(dim=1).item()
print('protos:', protos.numpy().round(4))
print('dists: ', dists.numpy().round(4))
print('pred: ', pred)| class | prototype | dist to query |
|---|---|---|
| 0 | [ 0.6177, 0.6338, 0.1551, -1.7046] | 1.9269 |
| 1 | [ 0.4111, -0.3812, -0.3202, 0.1425] | 0.6611 |
| 2 | [-0.8858, 1.1134, 0.1296, 0.3111] | 1.4621 |
Predicted class = 1 (nearest prototype, dist = 0.6611)
Why: The query at the origin is closest to class 1's prototype in L2 distance — the decision rule is argmin over distances, equivalent to argmax of the softmax above.
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Predicted class = 1 (nearest prototype, dist = 0.6611)
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:
Setup: 3 classes, 3-shot (9 support vectors in 4-dim embedding space). Compute each prototype and classify a query point at the origin.
Worked example
Use the pretrained backbone as f_φ. Build prototypes from 5-shot support; classify full test set (360 samples) by nearest prototype.
# use pretrained backbone as embedding function
pretrained.eval()
with torch.no_grad():
supp_emb = pretrained.features(torch.tensor(Xk)) # (50, 32)
query_emb = pretrained.features(torch.tensor(Xte)) # (360, 32)
# prototypes: mean embedding per class
classes = np.unique(yk)
protos = torch.stack([
supp_emb[yk == c].mean(dim=0) for c in classes
]) # shape (10, 32)
# classify: nearest prototype
dists = torch.cdist(query_emb, protos) # (360, 10)
preds = dists.argmin(dim=1).numpy()
acc = accuracy_score(yte, preds)
print(f'Prototypical net test acc: {acc:.4f}')| k-shot | proto net acc |
|---|---|
| 1 | 0.8111 |
| 3 | 0.9333 |
| 5 | 0.9389 |
| 10 | 0.9361 |
| 20 | 0.9417 |
Comparison
Comparison matrix
From Prototypical net on digits (5-shot): refill the proto net acc column from what you know. The rest of the table is as it appeared.
| k-shot | proto net acc |
|---|---|
| 1 | 0.8111 |
| 3 | 0.9333 |
| 5 | 0.9389 |
| 10 | 0.9361 |
| 20 | 0.9417 |
Anomaly
Predict first
A student writes this, and it looks reasonable:
We have raw pixel vectors (64-dim). Compute the mean pixel vector per class as the prototype and use Euclidean distance.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Raw pixel means are dominated by background intensity and pixel position, not semantic content.
Map samples through the learned embedding f_φ first. Compute prototypes in embedding space.
Why: Raw pixel means are dominated by background intensity and pixel position, not semantic content. Two very different digit-5 images average to a blurry blob that doesn't cleanly separate from digit-3. Distances in pixel space are not class-discriminative.
Trap
We have raw pixel vectors (64-dim). Compute the mean pixel vector per class as the prototype and use Euclidean distance.
proto[c] = mean(raw_pixels[class == c])
Why: Raw pixel means are dominated by background intensity and pixel position, not semantic content. Two very different digit-5 images average to a blurry blob that doesn't cleanly separate from digit-3. Distances in pixel space are not class-discriminative.
Map samples through the learned embedding f_φ first. Compute prototypes in embedding space.
proto[c] = mean(f_phi(x) for x in support if label == c)
Why: Embedding space is trained (or pretrained) to be class-discriminative. Euclidean distance in this space reliably separates classes — the entire point of metric learning.
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.
x is classified by computing its distance to each prototype and applying a softmax.; Zero-shot is only possible if the embedding space aligns visual and textual features — which is exactly what CLIP (Contrastive Language-Image Pretraining) trains.Adam(model.parameters(), lr=0.01) and train for 300 epochs.; We have raw pixel vectors (64-dim). Compute the mean pixel vector per class as the prototype and use Euclidean distance.Section
Part 3 of 4
Concept
Zero-shot learning classifies classes never seen during training by using side information — most commonly a textual description of the class.
| regime | support per class | how it works |
|---|---|---|
| supervised | unlimited | direct labels → train classifier |
| few-shot | k (small) | prototypes in embedding space |
| zero-shot | 0 | match query to class text descriptions |
Zero-shot is only possible if the embedding space aligns visual and textual features — which is exactly what CLIP (Contrastive Language-Image Pretraining) trains.
Explain it
Discussion prompt
Explain Zero-shot learning: no labeled examples at test time 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:
Zero-shot is only possible if the embedding space aligns visual and textual features — which is exactly what CLIP (Contrastive Language-Image Pretraining) trains.
Concept
CLIP trains an image encoder and a text encoder jointly so that matching (image, caption) pairs have high cosine similarity and non-matching pairs have low similarity — a contrastive objective over 400M image-text pairs.
\[ \text{sim}(v, t) = \frac{f_{\text{img}}(v)^\top f_{\text{txt}}(t)}{\|f_{\text{img}}(v)\|\,\|f_{\text{txt}}(t)\|} \]
At test time for a new class c, build a text prototype: f_txt("a photo of a c"). Classify image v by argmax over cosine similarity to all class text embeddings. No labeled images needed.
Analogy
Discussion prompt
Explain CLIP: aligning image and text embeddings 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:
At test time for a new class c, build a text prototype: f_txt("a photo of a c"). Classify image v by argmax over cosine similarity to all class text embeddings. No labeled images needed.
Concept
Transfer learning fails (or actively hurts) when the source and target domains are fundamentally incompatible.
The USAAIO exam may ask you to distinguish low source-target overlap (fine-tune carefully) from negative transfer (don't transfer at all).
Sorting
Sort into buckets
These are the pieces of Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods, out of order. Put each one back under the part of the lesson it belongs to.
Ranking
Put in order
These are the steps of Transfer learning decision recipe, scrambled. Put them back in order before the next slide shows you.
c_k = mean(f_φ(xᵢ)) per class; classify by argmin d(f_φ(x), c_k)f_txt("a photo of [class]") — no labeled images neededWhy: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.
Pattern
c_k = mean(f_φ(xᵢ)) per class; classify by argmin d(f_φ(x), c_k)f_txt("a photo of [class]") — no labeled images neededEdge cases
Discussion prompt
Transfer learning decision 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:
c_k = mean(f_φ(xᵢ)) per class; classify by argmin d(f_φ(x), c_k)f_txt("a photo of [class]") — no labeled images neededElimination
Eliminate the wrong options
30 labeled medical images, pretrained ImageNet backbone. Best fine-tuning strategy?
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: With only 30 labeled examples, training more parameters than you have data guarantees overfitting. The ImageNet backbone still encodes useful low-level features (edges, textures present in X-rays) — a linear probe has only 2048×3 parameters and can be safely fit on 30 samples.
Check
You have an ImageNet-pretrained ResNet and 30 labeled chest X-ray images (10 per class). Which strategy is most appropriate?
Check your understanding
30 labeled medical images, pretrained ImageNet backbone. Best fine-tuning strategy?
Answer: A
Why: With only 30 labeled examples, training more parameters than you have data guarantees overfitting. The ImageNet backbone still encodes useful low-level features (edges, textures present in X-rays) — a linear probe has only 2048×3 parameters and can be safely fit on 30 samples.
Prediction
Predict first
In a 10-way 5-shot prototypical network, classification of a query x proceeds by:
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: argmin_k d(f_φ(x), c_k) where c_k is the mean support embedding for class k
Why: The prototype c_k = mean of f_φ(xᵢ) over the support set for class k. Queries are assigned to the nearest prototype by Euclidean distance. The 5 distances from the digits trace verified this: query at origin → class 1 (dist 0.6611 < 1.4621 < 1.9269).
Check
Test your understanding of the prototype classification rule.
Check your understanding
In a 10-way 5-shot prototypical network, classification of a query x proceeds by:
Answer: A
Why: The prototype c_k = mean of f_φ(xᵢ) over the support set for class k. Queries are assigned to the nearest prototype by Euclidean distance. The 5 distances from the digits trace verified this: query at origin → class 1 (dist 0.6611 < 1.4621 < 1.9269).
Elimination
Eliminate the wrong options
CLIP enables zero-shot classification of a new class because:
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: CLIP's contrastive training aligns the image encoder and text encoder so matching (image, caption) pairs have high cosine similarity. At test time, f_txt('a photo of a [new class]') acts as a prototype in the shared space — no labeled images needed.
Check
Why does CLIP enable zero-shot classification?
Check your understanding
CLIP enables zero-shot classification of a new class because:
Answer: A
Why: CLIP's contrastive training aligns the image encoder and text encoder so matching (image, caption) pairs have high cosine similarity. At test time, f_txt('a photo of a [new class]') acts as a prototype in the shared space — no labeled images needed.
Section
Project
Concept
Pretrain a 2-layer net on digits (all 1437 training samples). Then use only 5 examples per class (50 total) and compare three strategies: from scratch, linear probe, prototypical network.
| # | task | expected output |
|---|---|---|
| 1 | Pretrain backbone (300 epochs) | test acc ≈ 0.967 |
| 2 | 5-shot linear probe | test acc ≈ 0.956 |
| 3 | 5-shot prototypical net | test acc ≈ 0.939 |
| 4 | Compare all three; explain the gap | table + 1-paragraph analysis |
Build rules: type every line; keep the features() method separate from forward() so the prototypical net can access embeddings without the classification head.
Counterexample
Discussion prompt
Pretrain a 2-layer net on digits (all 1437 training samples). Then use only 5 examples per class (50 total) and compare three strategies: from scratch, linear probe, prototypical network.
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; keep the features() method separate from forward() so the prototypical net can access embeddings without the classification head.
Worked example
Your turn: implement TinyNet (64→32 ReLU→10). Train on full digits training set for 300 epochs with Adam lr=0.01. What training loss do you expect at epoch 300?
Hint: F.cross_entropy(model(X_t), y_t) — call .backward() and .step() inside the loop.
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
digits = load_digits()
X = digits.data.astype('float32') / 16.0
y = digits.target
Xtr, Xte, ytr, yte = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y)
class TinyNet(nn.Module):
def __init__(self):
super().__init__()
self.backbone = nn.Sequential(nn.Linear(64, 32), nn.ReLU())
self.head = nn.Linear(32, 10)
def forward(self, x): return self.head(self.backbone(x))
def features(self, x): return self.backbone(x)
torch.manual_seed(0)
model = TinyNet()
opt = torch.optim.Adam(model.parameters(), lr=0.01)
Xtr_t = torch.tensor(Xtr); ytr_t = torch.tensor(ytr)
for ep in range(300):
opt.zero_grad()
F.cross_entropy(model(Xtr_t), ytr_t).backward(); opt.step()| epoch | train loss | test acc |
|---|---|---|
| 0 | 2.3273 | — |
| 50 | 0.1411 | — |
| 100 | 0.0575 | — |
| 300 | 0.0084 | 0.9667 |
Pattern
Step through it
Step through Milestone 1 — pretrain the backbone one row at a time. What is driving the change, and what would the row after the last one be?
Worked example
Your turn: sample 5 examples per class. Load the pretrained state_dict, freeze the backbone, train only model.head for 200 epochs with Adam lr=0.05. Predict whether it beats from-scratch.
Hint: for p in model.backbone.parameters(): p.requires_grad = False — then pass only model.head.parameters() to Adam.
def few_shot(X, y, k=5, seed=7):
rng = np.random.default_rng(seed)
idx = []
for c in np.unique(y):
idx.extend(rng.choice(np.where(y==c)[0], k, replace=False).tolist())
return X[idx], y[idx]
Xk, yk = few_shot(Xtr, ytr, k=5)
Xk_t = torch.tensor(Xk); yk_t = torch.tensor(yk)
torch.manual_seed(1)
probe = TinyNet()
probe.load_state_dict(model.state_dict()) # pretrained weights
for p in probe.backbone.parameters():
p.requires_grad = False
opt_p = torch.optim.Adam(probe.head.parameters(), lr=0.05)
for _ in range(200):
opt_p.zero_grad()
F.cross_entropy(probe(Xk_t), yk_t).backward(); opt_p.step()
probe.eval()
with torch.no_grad():
acc = accuracy_score(yte, probe(torch.tensor(Xte)).argmax(1).numpy())
print(f'linear probe 5-shot acc: {acc:.4f}')| strategy | test acc |
|---|---|
| from scratch (5-shot) | 0.8611 |
| linear probe (5-shot) | 0.9556 |
Trade off
Comparison matrix
From Milestone 2 — 5-shot linear probe: every row here is a choice with a cost. Fill the test acc column, then say which row you would actually pick and what you give up for it.
| strategy | test acc |
|---|---|
| from scratch (5-shot) | 0.8611 |
| linear probe (5-shot) | 0.9556 |
Worked example
Your turn: use model.features() to embed the 50 support examples and all 360 test queries. Compute class prototypes, compute pairwise L2 distances, classify by nearest prototype.
Hint: torch.cdist(query_emb, protos) gives an (n_queries, n_classes) distance matrix; .argmin(dim=1) classifies each query.
model.eval()
with torch.no_grad():
supp_emb = model.features(torch.tensor(Xk)) # (50, 32)
query_emb = model.features(torch.tensor(Xte)) # (360, 32)
classes = np.unique(yk)
protos = torch.stack([
supp_emb[yk == c].mean(dim=0) for c in classes
]) # (10, 32)
dists = torch.cdist(query_emb, protos) # (360, 10)
preds = dists.argmin(dim=1).numpy()
acc = accuracy_score(yte, preds)
print(f'proto net 5-shot acc: {acc:.4f}')| method | test acc | labeled support used |
|---|---|---|
| from scratch | 0.8611 | 50 |
| linear probe | 0.9556 | 50 |
| prototypical net | 0.9389 | 50 |
| full pretrain | 0.9667 | 1437 |
Discrimination
Sort into buckets
Sort these by labeled support used, from memory, without looking back at Milestone 3 — prototypical network. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Concept
# Full pipeline: pretrain → linear probe → proto net
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
digits = load_digits()
X = digits.data.astype('float32') / 16.0; y = digits.target
Xtr, Xte, ytr, yte = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y)
class TinyNet(nn.Module):
def __init__(self):
super().__init__()
self.backbone = nn.Sequential(nn.Linear(64,32), nn.ReLU())
self.head = nn.Linear(32,10)
def forward(self, x): return self.head(self.backbone(x))
def features(self, x): return self.backbone(x)
torch.manual_seed(0)
model = TinyNet()
opt = torch.optim.Adam(model.parameters(), lr=0.01)
Xtr_t=torch.tensor(Xtr); ytr_t=torch.tensor(ytr)
for _ in range(300):
opt.zero_grad(); F.cross_entropy(model(Xtr_t),ytr_t).backward(); opt.step()
def few_shot(X,y,k=5,seed=7):
rng=np.random.default_rng(seed); idx=[]
for c in np.unique(y): idx.extend(rng.choice(np.where(y==c)[0],k,replace=False).tolist())
return X[idx],y[idx]
Xk,yk=few_shot(Xtr,ytr,k=5); Xk_t=torch.tensor(Xk); yk_t=torch.tensor(yk)
# linear probe
torch.manual_seed(1)
probe=TinyNet(); probe.load_state_dict(model.state_dict())
for p in probe.backbone.parameters(): p.requires_grad=False
opt_p=torch.optim.Adam(probe.head.parameters(),lr=0.05)
for _ in range(200):
opt_p.zero_grad(); F.cross_entropy(probe(Xk_t),yk_t).backward(); opt_p.step()
# proto net
model.eval()
with torch.no_grad():
se=model.features(Xk_t); qe=model.features(torch.tensor(Xte))
protos=torch.stack([se[yk==c].mean(0) for c in np.unique(yk)])
preds=torch.cdist(qe,protos).argmin(1).numpy()
print('proto acc:', accuracy_score(yte,preds))| output | value |
|---|---|
| linear probe acc | 0.9556 |
| proto net acc | 0.9389 |
If your linear probe beats from-scratch (0.86) and your proto net reaches 0.93 with no head training at all — you've demonstrated the core value of pretrained representations for low-data regimes.
Comparison
Comparison matrix
From The full program: refill the value column from what you know. The rest of the table is as it appeared.
| output | value |
|---|---|
| linear probe acc | 0.9556 |
| proto net acc | 0.9389 |
Concept
Out loud, slides closed: (1) Given a dataset of 20 medical images per class and an ImageNet backbone, what sequence of steps do you take and why? (2) Explain the prototypical network classification rule in one sentence without notation. (3) Why does CLIP enable zero-shot and a standard ImageNet classifier does not?
Stretch (homework): implement BERT fine-tuning on a text classification task using HuggingFace transformers; compare training from scratch to fine-tuning the pretrained BERT. Also: explore prototypical networks on the Omniglot dataset as described in the original Snell et al. 2017 paper.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Why transfer? What the backbone knows · Few-shot learning & prototypical networks · Zero-shot learning & CLIP · Your turn: fine-tune & few-shot. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
| concept | the one thing to remember |
|---|---|
| linear probe | freeze backbone; only new head is trained — safe at k ≈ 5-50 |
| gradual unfreeze | head first (large lr), then backbone (small lr) — avoids forgetting |
| prototypical net | c_k = mean(f_φ(x)) over support; classify by argmin distance |
| zero-shot (CLIP) | text description replaces labeled images as class prototype |
| when NOT to | extreme domain gap → train from scratch; no negative transfer |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.