Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods

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

What this lesson covers

The lesson, slide by slide

1. Transfer Learning, Fine-Tuning & Few-Shot

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.

2. By the end of this lesson you can

Objectives

  1. Choose between linear probe, gradual unfreeze, and full fine-tune for a given data regime
  2. Explain domain adaptation and identify when distribution shift makes transfer unsafe
  3. Implement k-shot prototypical networks: compute class prototypes, classify by nearest prototype
  4. Describe zero-shot learning (CLIP) and explain why it is the limiting case of few-shot
  5. Diagnose when NOT to use transfer learning (highly mismatched source/target domains)

3. What survived from Learning Rate Schedules?

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.

4. Why transfer? What the backbone knows

Section

Part 1 of 4

5. The pretrained backbone as a feature extractor

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 typewhat it encodesreusable?
early backbonelow-level features (edges, curves)almost always
deep backbonesemantic concepts (dog, wheel)usually, with care
head (linear)class boundaries for source taskno — 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.

6. Fill in: what it encodes for The pretrained backbone as a feature…

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 typewhat it encodesreusable?
early backbonelow-level features (edges, curves)almost always
deep backbonesemantic concepts (dog, wheel)usually, with care
head (linear)class boundaries for source taskno — replace it

7. Domain adaptation: source → target

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 → targetshift severitystrategy
ImageNet → pet photoslowlinear probe usually sufficient
ImageNet → medical imagesmediumgradual unfreeze + domain data
natural images → satellitehighfull fine-tune with larger dataset
vision → languageextremedo NOT transfer; train from scratch

8. What each one costs: Domain adaptation: source → target

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 → targetshift severitystrategy
ImageNet → pet photoslowlinear probe usually sufficient
ImageNet → medical imagesmediumgradual unfreeze + domain data
natural images → satellitehighfull fine-tune with larger dataset
vision → languageextremedo NOT transfer; train from scratch

9. Why linear probe first?

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.

10. Break it if you can: Why linear probe first?

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.

11. Linear probe vs. from-scratch on 5-shot

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()
epochtraining loss
02.3273
500.1411
1000.0575
2000.0185
3000.0084

12. Watch it run: Linear probe vs. from-scratch on 5-shot

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?

  1. Step 1: epoch is 0
  2. Step 2: epoch is 50
  3. Step 3: epoch is 100
  4. Step 4: epoch is 200
  5. Step 5: epoch is 300

13. 5-shot: probe vs. scratch comparison

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()
strategytest 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

14. Something is wrong here: full fine-tune before stabilizing the head

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.

15. Trap: full fine-tune before stabilizing the head

Trap

The 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.

The fix

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.

16. Break it on purpose: full fine-tune before stabilizing the head

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.

17. Few-shot learning & prototypical networks

Section

Part 2 of 4

18. The few-shot setup

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.

termmeaningexample
N-waynumber of classes10-way = 10 classes
k-shotexamples per class in support5-shot = 5 per class
support setlabeled examples provided at test5-shot × 10 = 50 samples
query setunlabeled samples to classifythe actual test queries
episodeone (support, query) pair for trainingmeta-learning iterate

19. Prototypical networks: class prototypes

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'}))} \]

20. By analogy: Prototypical networks: class prototypes

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).

21. Guess the shape of the answer: Prototype computation — explicit trace

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.

22. Prototype computation — explicit trace

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)
classprototypedist 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.

23. Work backwards from the answer: Prototype computation — explicit trace

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.

24. Prototypical net on digits (5-shot)

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-shotproto net acc
10.8111
30.9333
50.9389
100.9361
200.9417

25. Fill in: proto net acc for Prototypical net on digits (5-shot)

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-shotproto net acc
10.8111
30.9333
50.9389
100.9361
200.9417

26. Something is wrong here: prototype = centroid in input space (not embedding…

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.

27. Trap: prototype = centroid in input space (not embedding space)

Trap

The 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.

The fix

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.

28. Which of these survive contact with Lesson 63: Transfer Learning, Fine-Tuning &…?

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.

Holds up
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 query 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.
Breaks
We have a pretrained backbone and 50 labeled samples. Jump straight to 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.
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

29. Zero-shot learning & CLIP

Section

Part 3 of 4

30. Zero-shot learning: no labeled examples at test time

Concept

Zero-shot learning classifies classes never seen during training by using side information — most commonly a textual description of the class.

regimesupport per classhow it works
supervisedunlimiteddirect labels → train classifier
few-shotk (small)prototypes in embedding space
zero-shot0match 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.

31. Teach it back: Zero-shot learning: no labeled examples at test time

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.

32. CLIP: aligning image and text embeddings

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.

33. By analogy: CLIP: aligning image and text embeddings

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.

34. When NOT to use transfer learning

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).

35. Where does each piece belong: Lesson 63: Transfer Learning, Fine-Tuning &…

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.

Why transfer? What the backbone knows
The pretrained backbone as a feature extractor; Domain adaptation: source → target; Why linear probe first?
Few-shot learning & prototypical networks
The few-shot setup; Prototypical networks: class prototypes; Prototype computation — explicit trace
Zero-shot learning & CLIP
Zero-shot learning: no labeled examples at test time; CLIP: aligning image and text embeddings; When NOT to use transfer learning
s1
Why transfer? What the backbone knows is where Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods puts The pretrained backbone as a feature extractor, Domain adaptation: source → target, Why linear probe first?. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
Few-shot learning & prototypical networks is where Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods puts The few-shot setup, Prototypical networks: class prototypes, Prototype computation — explicit trace. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
Zero-shot learning & CLIP is where Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods puts Zero-shot learning: no labeled examples at test time, CLIP: aligning image and text embeddings, When NOT to use transfer learning. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

36. Rebuild the recipe: Transfer learning decision recipe

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.

  1. Assess domain gap: low gap → linear probe; medium → gradual unfreeze; extreme → train from scratch
  2. Linear probe first: freeze backbone, train new head; fast sanity-check of pretrained features
  3. Gradual unfreeze: stage 1 head (lr ≈ 1e-2), stage 2 full (lr ≈ 1e-3); prevents catastrophic forgetting
  4. Few-shot → prototypical net: compute c_k = mean(f_φ(xᵢ)) per class; classify by argmin d(f_φ(x), c_k)
  5. Zero-shot (CLIP): replace image prototypes with f_txt("a photo of [class]") — no labeled images needed

Why: 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.

37. Transfer learning decision recipe

Pattern

  1. Assess domain gap: low gap → linear probe; medium → gradual unfreeze; extreme → train from scratch
  2. Linear probe first: freeze backbone, train new head; fast sanity-check of pretrained features
  3. Gradual unfreeze: stage 1 head (lr ≈ 1e-2), stage 2 full (lr ≈ 1e-3); prevents catastrophic forgetting
  4. Few-shot → prototypical net: compute c_k = mean(f_φ(xᵢ)) per class; classify by argmin d(f_φ(x), c_k)
  5. Zero-shot (CLIP): replace image prototypes with f_txt("a photo of [class]") — no labeled images needed

38. Where does it stop working: Transfer learning decision recipe

Edge 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:

  1. Assess domain gap: low gap → linear probe; medium → gradual unfreeze; extreme → train from scratch
  2. Linear probe first: freeze backbone, train new head; fast sanity-check of pretrained features
  3. Gradual unfreeze: stage 1 head (lr ≈ 1e-2), stage 2 full (lr ≈ 1e-3); prevents catastrophic forgetting
  4. Few-shot → prototypical net: compute c_k = mean(f_φ(xᵢ)) per class; classify by argmin d(f_φ(x), c_k)
  5. Zero-shot (CLIP): replace image prototypes with f_txt("a photo of [class]") — no labeled images needed

39. Rule out three: Check yourself — fine-tuning strategy

Elimination

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.

  • A. Linear probe (freeze backbone, train head only)
  • B. Full fine-tune immediately at lr=0.01
  • C. Train a new ResNet from scratch on the 30 images
  • D. Zero-shot using CLIP text prototypes

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.

40. Check yourself — fine-tuning strategy

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?

  • A. Linear probe (freeze backbone, train head only) (correct)
  • B. Full fine-tune immediately at lr=0.01
  • C. Train a new ResNet from scratch on the 30 images
  • D. Zero-shot using CLIP text prototypes

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.

Why B tempts people
Full fine-tune at a large lr with 30 samples will overwrite the pretrained backbone (catastrophic forgetting) and overfit the tiny dataset. This is the textbook case FOR freezing.
Why C tempts people
Training from scratch requires vastly more data. 30 samples with ~25M ResNet parameters will achieve near-zero training loss and near-chance test accuracy.
Why D tempts people
CLIP zero-shot works for natural-image categories described by text. Medical imaging classes ('pleural effusion', 'pneumothorax') are not reliably captured by generic CLIP text prototypes without domain-specific fine-tuning.

41. Answer it before you see the options: Check yourself — prototypical networks

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).

42. Check yourself — prototypical networks

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:

  • A. argmin_k d(f_φ(x), c_k) where c_k is the mean support embedding for class k (correct)
  • B. argmax_k f_φ(x)ᵀ w_k where w_k is a trained linear weight vector for class k
  • C. argmin_k d(x, c_k) where c_k is the mean raw input for class k
  • D. argmax_k p(y=k) · p(x | y=k) where p(x|y=k) is Gaussian

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).

Why B tempts people
That is a standard linear classifier head — not metric-based. It requires training class weights w_k; prototypical nets require no learned parameters beyond f_φ.
Why C tempts people
Using raw input x (not embedded) fails because the input space is not discriminative — prototype distance in pixel/raw space gives poor separation.
Why D tempts people
That is Naive Bayes (Gaussian class-conditional). Prototypical nets do not model class-conditional densities; they use Euclidean distance in embedding space.

43. Rule out three: Check yourself — zero-shot CLIP

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.

  • A. Its shared embedding space aligns image and text vectors, so a text description serves as a class prototype
  • B. It was trained on every possible image class so all classes are seen during training
  • C. Its image encoder outputs one-hot vectors that match class indices
  • D. Zero-shot works by kNN over the full training set of 400M images

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.

44. Check yourself — zero-shot CLIP

Check

Why does CLIP enable zero-shot classification?

Check your understanding

CLIP enables zero-shot classification of a new class because:

  • A. Its shared embedding space aligns image and text vectors, so a text description serves as a class prototype (correct)
  • B. It was trained on every possible image class so all classes are seen during training
  • C. Its image encoder outputs one-hot vectors that match class indices
  • D. Zero-shot works by kNN over the full training set of 400M images

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.

Why B tempts people
CLIP doesn't memorize every class — it learns a general alignment. It generalizes to classes described but never explicitly labeled in training because the alignment is over semantic content, not discrete labels.
Why C tempts people
CLIP encoders produce dense continuous embeddings, not one-hot vectors. The classification is done by cosine similarity, not index lookup.
Why D tempts people
kNN over 400M training images would be computationally prohibitive and would still require image-level labels for each example, not descriptions.

45. Your turn: fine-tune & few-shot

Section

Project

46. Project: transfer learning on digits

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.

#taskexpected output
1Pretrain backbone (300 epochs)test acc ≈ 0.967
25-shot linear probetest acc ≈ 0.956
35-shot prototypical nettest acc ≈ 0.939
4Compare all three; explain the gaptable + 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.

47. Break it if you can: Project: transfer learning on digits

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.

48. Milestone 1 — pretrain the backbone

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()
epochtrain losstest acc
02.3273—
500.1411—
1000.0575—
3000.00840.9667

49. Watch it run: Milestone 1 — pretrain the backbone

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?

  1. Step 1: epoch is 0
  2. Step 2: epoch is 50
  3. Step 3: epoch is 100
  4. Step 4: epoch is 300

50. Milestone 2 — 5-shot linear probe

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}')
strategytest acc
from scratch (5-shot)0.8611
linear probe (5-shot)0.9556

51. What each one costs: Milestone 2 — 5-shot linear probe

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.

strategytest acc
from scratch (5-shot)0.8611
linear probe (5-shot)0.9556

52. Milestone 3 — prototypical network

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}')
methodtest acclabeled support used
from scratch0.861150
linear probe0.955650
prototypical net0.938950
full pretrain0.96671437

53. Which is which, by labeled support used

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.

50
from scratch; linear probe; prototypical net
1437
full pretrain
g1
labeled support used is "50" for from scratch, linear probe, prototypical net — that is what the table on "Milestone 3 — prototypical network" records, and it is the single property separating this group from the rest.
g2
labeled support used is "1437" for full pretrain — that is what the table on "Milestone 3 — prototypical network" records, and it is the single property separating this group from the rest.

54. The full program

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))
outputvalue
linear probe acc0.9556
proto net acc0.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.

55. Fill in: value for The full program

Comparison

Comparison matrix

From The full program: refill the value column from what you know. The rest of the table is as it appeared.

outputvalue
linear probe acc0.9556
proto net acc0.9389

56. Show it off

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.

57. Connect it up: Lesson 63: Transfer Learning, Fine-Tuning & Few-Shot Methods

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.

58. What you can do now

Recap

conceptthe one thing to remember
linear probefreeze backbone; only new head is trained — safe at k ≈ 5-50
gradual unfreezehead first (large lr), then backbone (small lr) — avoids forgetting
prototypical netc_k = mean(f_φ(x)) over support; classify by argmin distance
zero-shot (CLIP)text description replaces labeled images as class prototype
when NOT toextreme domain gap → train from scratch; no negative transfer

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 63 (Week 22 — Fine-tuning, Few-Shot, Prototypical Networks) — Barron · USAAIO Round 2 Preparation, 2026
  2. Transfer learning strategies and prototypical network accuracies verified on load_digits — torch 2.7.1 + sklearn, real execution, June 2026

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108