Course Content
Transfer Learning and Pretraining
3 sections · 7 lessons
Fine-Tune a ResNet on a Real Image Dataset
Two people build the same flower classifier on the same public dataset with the same ResNet-50. One reports 94.6% test accuracy. The other reports 98.9% and is delighted, right up until someone asks how many images they trained on.
The answer is 6,149. The published benchmark trains on 1,020.
Oxford Flowers-102 has an unusual property: its test split is six times larger than its training split. The dataset was designed for the low-data regime — 10 images per class for training, 10 for validation, and everything else held out. Anyone who assumes the biggest split must be the training set trains on six times the intended data, evaluates on the intended training set, and produces a number that cannot be compared with any published result.
That is the shape of most failures in a project like this. Not exotic modelling errors — a wrong split, an augmentation applied to the validation set, a normalisation constant copied from the wrong recipe. This project is built so you meet each of those deliberately.
You will take a pretrained ResNet-50 and adapt it to a fine-grained classification problem, in four measured stages, ending with an honest error analysis rather than a single accuracy number. Budget six to eight hours. The compute fits comfortably on a free Colab GPU.
The project at a glance
| Stage | What you build | Time | Target result |
|---|---|---|---|
| 1. Data preparation | Loaders, split verification, augmentation sized to the data | ~1 hour | Correct splits, verified visually |
| 2. Model construction | Pretrained backbone, replaced head, reusable train/eval loop | ~1 hour | Forward pass with the right output shape |
| 3. Frozen baseline | Feature extraction, backbone frozen | ~2 hours | Roughly 80–87% test accuracy |
| 4. Progressive fine-tuning | Staged unfreezing, discriminative learning rates | ~2 hours | Roughly 90–95% test accuracy |
| 5. Evaluation and analysis | Confusion matrix, per-class metrics, error inspection | ~1 hour | A written account of what the model gets wrong |
Dataset options
| Dataset | Classes | Train images | Difficulty | Why choose it |
|---|---|---|---|---|
| Oxford Flowers-102 | 102 | 1,020 | Hard | Genuinely small data — 10 per class. The recommended default. |
| Oxford-IIIT Pet | 37 | 3,680 | Medium | Breed-level fine-grained, close to ImageNet, forgiving |
| Stanford Cars | 196 | 8,144 | Hard | Very fine-grained; model year matters more than shape |
| Food-101 | 101 | 75,750 | Medium | Large enough that full fine-tuning genuinely wins |
| Caltech-101 | 101 | ~3,000 | Easy | Fastest to iterate on; heavily unbalanced classes |
Flowers-102 is the recommended choice precisely because 10 images per class puts you in the regime where every decision in this project matters. With 25.6 million parameters and 1,020 training images you have about 25,000 free parameters per example — full fine-tuning at a careless learning rate will lose to doing nothing at all.
flower-transfer/├── data/ # downloaded automatically, do not commit├── src/│ ├── data.py # datasets, transforms, loaders│ ├── model.py # backbone loading, head replacement, freezing│ ├── engine.py # train_one_epoch, evaluate, early stopping│ └── analyse.py # confusion matrix, per-class metrics├── checkpoints/│ ├── phase1_frozen.pt│ ├── phase2_layer4.pt│ └── phase3_discriminative.pt├── results/│ ├── history.csv # per-epoch loss and accuracy for every phase│ └── confusion.png└── report.mdPart 1: data preparation
Loading the data, and checking the split
1from torchvision import datasets2from torchvision.models import ResNet50_Weights34weights = ResNet50_Weights.IMAGENET1K_V256train_raw = datasets.Flowers102(root="data", split="train", download=True)7val_raw = datasets.Flowers102(root="data", split="val", download=True)8test_raw = datasets.Flowers102(root="data", split="test", download=True)910print(len(train_raw), len(val_raw), len(test_raw)) # 1020 1020 6149Print those three numbers and look at them. If the test set is larger than the training set, that is correct for this dataset and you must not "fix" it. Resisting the urge to swap them is the whole point of this step.
Two more checks before going further. Confirm the label range is what you expect — Flowers-102 uses 0–101, and some torchvision datasets have historically used 1-indexed labels, which produces an off-by-one that shows up as a plausible-but-wrong accuracy. And confirm no image appears in more than one split; on datasets you assemble yourself, hash the file contents and check for duplicates across splits, because near-duplicate leakage is the most common way people accidentally report 99%.
Explore before you model
1from collections import Counter2import numpy as np34counts = Counter(label for _, label in train_raw)5c = np.array(sorted(counts.values()))6print(f"classes {len(counts)} min {c.min()} max {c.max()} median {np.median(c)}")78sizes = [img.size for img, _ in list(train_raw)[:200]]9ws, hs = zip(*sizes)10print(f"width {min(ws)}-{max(ws)}, height {min(hs)}-{max(hs)}")You are looking for three things. Class balance: Flowers-102 is exactly 10 per class in training, so no imbalance handling is needed — on Caltech-101 you would find roughly a 20:1 ratio between the largest and smallest classes and would need class weights. Image sizes: highly variable aspect ratios mean Resize then CenterCrop may cut off the subject. What the images actually look like: display a 5×5 grid and look at them, because half an hour of looking will tell you which augmentations are safe.
Augmentation, sized to 10 images per class
1from torchvision import transforms23norm = transforms.Normalize(mean=[0.485, 0.456, 0.406],4 std=[0.229, 0.224, 0.225])56train_tf = transforms.Compose([7 transforms.RandomResizedCrop(224, scale=(0.5, 1.0), ratio=(0.75, 1.33)),8 transforms.RandomHorizontalFlip(),9 transforms.RandomVerticalFlip(), # flowers have no fixed "up"10 transforms.RandomRotation(30),11 transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.05),12 transforms.ToTensor(), norm,13 transforms.RandomErasing(p=0.25, scale=(0.02, 0.2)),14])1516eval_tf = transforms.Compose([17 transforms.Resize(232), # matches the V2 weights recipe18 transforms.CenterCrop(224),19 transforms.ToTensor(), norm,20])Three things in there are deliberate. The crop scale goes down to 0.5 — aggressive, justified by having only 10 images per class. Vertical flips are enabled because a flower photographed from above has no canonical orientation; on a dataset of cars or people this would be actively harmful. And hue jitter is kept small at 0.05, because on flowers colour is a genuine class signal and shifting hue by a large amount can turn one species into another.
The evaluation transform uses Resize(232), not 256, because that is what the IMAGENET1K_V2 recipe used. Rather than trusting that from memory, read it off the weights object with weights.transforms() and match it.
Augmentation belongs to the training split only. Applying random crops or flips at evaluation time makes your validation number noisy and optimistic in unpredictable directions, and it is one of the easiest mistakes to make when a single dataset object is shared between loaders.
Wrapping a dataset so two transforms can share it
Torchvision datasets hold one transform, set at construction. If you want the same underlying images with different transforms for training and evaluation, wrap them.
1from torch.utils.data import Dataset, DataLoader23class TransformedDataset(Dataset):4 def __init__(self, base, transform):5 self.base, self.transform = base, transform6 def __len__(self):7 return len(self.base)8 def __getitem__(self, i):9 img, label = self.base[i]10 return self.transform(img), label1112train_ds = TransformedDataset(train_raw, train_tf)13val_ds = TransformedDataset(val_raw, eval_tf)14test_ds = TransformedDataset(test_raw, eval_tf)1516train_dl = DataLoader(train_ds, batch_size=32, shuffle=True,17 num_workers=4, pin_memory=True, drop_last=True)18val_dl = DataLoader(val_ds, batch_size=64, num_workers=4, pin_memory=True)19test_dl = DataLoader(test_ds, batch_size=64, num_workers=4, pin_memory=True)drop_last=True on the training loader avoids a final batch of size 1, which crashes BatchNorm — with 1,020 images and batch size 32 you get 31 full batches and a remainder of 28, so it is safe here, but the habit costs nothing.
Part 2: building the model
1import torch, torch.nn as nn2from torchvision import models34def build_model(num_classes=102, dropout=0.4):5 weights = models.ResNet50_Weights.IMAGENET1K_V26 model = models.resnet50(weights=weights)78 for p in model.parameters(): # freeze FIRST9 p.requires_grad = False1011 model.fc = nn.Sequential( # then replace - new layers train12 nn.Dropout(dropout),13 nn.Linear(model.fc.in_features, num_classes),14 )15 return model1617model = build_model().cuda()1819total = sum(p.numel() for p in model.parameters())20trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)21print(f"{trainable:,} / {total:,} trainable ({100*trainable/total:.3f}%)")22# 208,998 / 23,717,030 trainable (0.881%)2324with torch.no_grad():25 print(model(torch.randn(2, 3, 224, 224).cuda()).shape) # [2, 102]Verify both printed lines. The head is 2048×102+102=208,998 parameters — with 1,020 training images that is 205 parameters per example, which is a regime a small dataset can actually support. If the trainable percentage reads 100%, your freeze loop ran after the head replacement or not at all; if the output's second dimension is 1000, you never replaced the head.
Note also the deprecated API. Old tutorials write models.resnet50(pretrained=True), which current torchvision still accepts but warns about. The weights= enum replaced it because a boolean cannot express which ImageNet weights you want — IMAGENET1K_V1 scores 76.1% top-1 and IMAGENET1K_V2 scores 80.9% — and because the enum carries the correct preprocessing with it.
A training loop you will reuse four times
1def run_epoch(model, loader, criterion, optimiser=None, device="cuda"):2 training = optimiser is not None3 model.train(training)4 total_loss, correct, n = 0.0, 0, 05 for x, y in loader:6 x, y = x.to(device, non_blocking=True), y.to(device)7 with torch.set_grad_enabled(training):8 out = model(x)9 loss = criterion(out, y)10 if training:11 optimiser.zero_grad(set_to_none=True)12 loss.backward()13 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)14 optimiser.step()15 total_loss += loss.item() * y.size(0)16 correct += (out.argmax(1) == y).sum().item()17 n += y.size(0)18 return total_loss / n, correct / n192021def train_phase(model, name, optimiser, epochs, criterion,22 train_dl, val_dl, patience=8):23 scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimiser, T_max=epochs)24 best, bad = 0.0, 025 for epoch in range(epochs):26 tr_loss, tr_acc = run_epoch(model, train_dl, criterion, optimiser)27 va_loss, va_acc = run_epoch(model, val_dl, criterion)28 scheduler.step()29 print(f"[{name}] {epoch:2d} train {tr_acc:.4f} val {va_acc:.4f} "30 f"gap {tr_acc - va_acc:+.4f}")31 if va_acc > best:32 best, bad = va_acc, 033 torch.save(model.state_dict(), f"checkpoints/{name}.pt")34 else:35 bad += 136 if bad >= patience:37 break38 model.load_state_dict(torch.load(f"checkpoints/{name}.pt"))39 return bestPrint the train–validation gap every epoch, not just the two accuracies. It is the number that tells you what to do next: under 5 points means you have headroom to unfreeze more, over 15 means you are memorising and should regularise instead.
Part 3: the frozen baseline
1criterion = nn.CrossEntropyLoss(label_smoothing=0.1)2optimiser = torch.optim.AdamW(3 [p for p in model.parameters() if p.requires_grad],4 lr=1e-3, weight_decay=1e-4)56best_frozen = train_phase(model, "phase1_frozen", optimiser,7 epochs=30, criterion=criterion,8 train_dl=train_dl, val_dl=val_dl)9print(f"frozen baseline: {best_frozen:.4f}") # validation; expect mid-0.80sThis should take ten to fifteen minutes and land in the mid-80s on validation; one run with exactly these settings gave 84% validation and 81% test accuracy. (If you only ever train a head, the older IMAGENET1K_V1 weights give a few points more here — see the note on V1 and V2 in the backbone lesson.) That number is now the floor for the entire project — every later phase must beat it, and any phase that does not gets discarded rather than explained away.
If it lands near 1% — chance level for 102 classes — stop and debug the data pipeline. The usual causes are labels shuffled relative to images, normalisation constants that do not match the weights, or a freeze loop that ran in the wrong order. No amount of fine-tuning fixes a broken loader.
Part 4: progressive fine-tuning
Now let the backbone move, in controlled stages. The rule that prevents disaster: the head must be trained before any backbone weight is unfrozen. A randomly initialised head produces large, meaningless gradients, and at a from-scratch learning rate those gradients destroy the pretrained features in the first fifty steps — the loss spikes, then settles at chance, and never recovers.
1def set_trainable(model, groups):2 for p in model.parameters():3 p.requires_grad = False4 for g in groups:5 for p in g.parameters():6 p.requires_grad = True78# Phase 2: add layer4 only, at a 10x lower learning rate.9set_trainable(model, [model.fc, model.layer4])10opt2 = torch.optim.AdamW(11 [p for p in model.parameters() if p.requires_grad],12 lr=1e-4, weight_decay=1e-4)13best_p2 = train_phase(model, "phase2_layer4", opt2, epochs=25,14 criterion=criterion, train_dl=train_dl, val_dl=val_dl)15print(f"phase 2: {best_p2:.4f}") # expect ~0.92Then discriminative learning rates across the whole network. Every layer trains, but deeper layers get exponentially smaller rates — early layers detect edges and textures that are as valid for flowers as for ImageNet and need almost no change, while late layers encode whole-object ImageNet concepts and need the most.
1for p in model.parameters():2 p.requires_grad = True34base_lr, factor = 1e-4, 2.65groups = [model.fc, model.layer4, model.layer3, model.layer2, model.layer1]6opt3 = torch.optim.AdamW(7 [{"params": g.parameters(), "lr": base_lr / factor ** d}8 for d, g in enumerate(groups)],9 weight_decay=1e-4)1011best_p3 = train_phase(model, "phase3_discriminative", opt3, epochs=25,12 criterion=criterion, train_dl=train_dl, val_dl=val_dl)13print(f"phase 3: {best_p3:.4f}") # expect ~0.94| Group | Formula | Learning rate | Relative to head |
|---|---|---|---|
fc | 10−4/2.60 | 1.00e-4 | 1× |
layer4 | 10−4/2.61 | 3.85e-5 | 0.38× |
layer3 | 10−4/2.62 | 1.48e-5 | 0.15× |
layer2 | 10−4/2.63 | 5.69e-6 | 0.057× |
layer1 | 10−4/2.64 | 2.19e-6 | 0.022× |
The bottom of the network moves 46 times more slowly than the top. For scale, the same run as above reached 92.5% validation (90.3% test) after phase 2 and 94.0% validation (91.7% test) after phase 3 — test accuracy sits a few points below validation because validation was used to pick the checkpoint, which makes it slightly optimistic. Validate between phases and keep the best checkpoint, not the last phase — on a 10-images-per-class dataset it is entirely possible that phase 3 loses to phase 2, and that result is information, not a failure. It tells you the dataset cannot support that many free parameters and that further effort belongs in the data.
Part 5: evaluation and error analysis
Run the test set exactly once, with the best checkpoint, and then spend your time understanding the errors rather than chasing the number.
1import numpy as np2from sklearn.metrics import confusion_matrix, classification_report34@torch.no_grad()5def predict(model, loader, device="cuda"):6 model.eval()7 preds, targets, confs = [], [], []8 for x, y in loader:9 probs = model(x.to(device)).softmax(1)10 conf, pred = probs.max(1)11 preds.append(pred.cpu()); targets.append(y); confs.append(conf.cpu())12 return (torch.cat(preds).numpy(), torch.cat(targets).numpy(),13 torch.cat(confs).numpy())1415preds, targets, confs = predict(model, test_dl)16cm = confusion_matrix(targets, preds)17print(classification_report(targets, preds, digits=3))1819# The most confused pairs, ignoring the diagonal.20off = cm.copy(); np.fill_diagonal(off, 0)21for _ in range(10):22 i, j = np.unravel_index(off.argmax(), off.shape)23 print(f"true {i:3d} -> predicted {j:3d} : {off[i, j]} times")24 off[i, j] = 0Reading a confusion matrix
A 102×102 matrix is unreadable as a picture, which is why the "most confused pairs" listing above is more useful than a heatmap. But the reading skill is best learned on something small. Here is a four-class extract, rows being the true class and columns the prediction:
| True ↓ / Pred → | Sunflower | Daisy | Coneflower | Black-eyed Susan | Total |
|---|---|---|---|---|---|
| Sunflower | 58 | 1 | 4 | 2 | 65 |
| Daisy | 2 | 60 | 1 | 2 | 65 |
| Coneflower | 5 | 0 | 47 | 13 | 65 |
| Black-eyed Susan | 3 | 1 | 16 | 45 | 65 |
Overall accuracy is (58+60+47+45)/260=210/260=80.8%, which sounds acceptable and hides everything interesting.
Now read down the Coneflower column. It contains 4+1+47+16=68 predictions, of which 47 are correct, so precision is 47/68=69.1%. Read across the Coneflower row: 65 true examples, 47 found, so recall is 47/65=72.3%, giving an F1 of 70.7%. Compare that with Daisy, whose precision is 60/62=96.8% and recall 60/65=92.3%.
The dominant error is the Coneflower–Black-eyed Susan pair: 13 one way and 16 the other, so 29 of the 130 examples in those two classes are confused with each other — 22%. Every other cell is single digits. This is a specific, actionable finding, and it is invisible in the 80.8% headline.
It also has an obvious explanation once you look at the images: both are daisy-shaped with a raised dark central cone, differing mainly in petal colour saturation and cone shape. Which immediately implicates the aggressive ColorJitter(saturation=0.3) in the training transform, since saturation is one of the few features separating them.
A single accuracy number tells you how well you did. A confusion matrix tells you what to do next. Any project that reports only the first has stopped one step short of the useful part.
Look at the actual errors
1wrong = np.where(preds != targets)[0]2order = wrong[np.argsort(-confs[wrong])] # most confident mistakes first3for i in order[:20]:4 print(f"true {targets[i]:3d} pred {preds[i]:3d} confidence {confs[i]:.3f}")Display those twenty images and sort them into three buckets: label errors (the dataset is wrong — common, and you should count how many), genuinely ambiguous (two species that look alike, or a photo where the flower is barely visible), and real model failures (a clear image the model should have got). Only the third bucket is worth modelling effort. If half your errors are label noise, your effective ceiling is lower than 100% and chasing the last two points is wasted work.
Deliverables
| Deliverable | What it must contain |
|---|---|
| Working code | Runs end to end from a clean checkout; fixed random seeds; the split sizes printed and unaltered |
| Training history | Per-epoch train and validation accuracy for all three phases, in one plot with phase boundaries marked |
| Results table | Test accuracy and macro-F1 for each phase, so the gain from each stage is visible |
| Confusion analysis | The ten most confused class pairs, with images, and a hypothesis for each |
| Error inspection | Twenty highest-confidence errors, categorised as label noise / ambiguous / model failure |
| Written report | What you tried, what failed and why, what you would do with another week |
| Criterion | Weight | What full marks looks like |
|---|---|---|
| Correctness of the pipeline | 25% | Splits respected, no leakage, augmentation on training only, preprocessing matched to the weights |
| Staged methodology | 25% | Frozen baseline established first; each phase validated before the next; losing phases discarded |
| Final performance | 20% | Test accuracy above the frozen baseline by a clear margin |
| Error analysis | 20% | Specific, image-level findings with plausible causes — not a restated accuracy |
| Reproducibility | 10% | Seeds fixed, environment pinned, results reproduced within a point on a second run |
Optional extensions
| Extension | What to do | Realistic gain |
|---|---|---|
| Ensembling | Average the softmax outputs of models from different seeds or backbones. Also try test-time augmentation: average predictions over the original plus a horizontal flip. | +1 to +3 points, at a proportional inference cost |
| Knowledge distillation | Train a MobileNetV3 student on the ResNet-50's soft outputs, using a temperature of 3–5 on both. The soft probabilities carry the near-miss structure that hard labels discard. | Student recovers most of the teacher's accuracy at ~5× the speed |
| Deployment | Export to ONNX (or with torch.export; TorchScript is deprecated), apply 8-bit static post-training quantisation, and measure latency and size before and after. Dynamic quantisation only touches Linear layers, which are a small part of a ResNet. | ~4× smaller, 2–3× faster on CPU, typically under a point lost |
| Explanation | Run Grad-CAM on the confused pairs to see which pixels drove each decision. Check whether the model looks at the flower or at the background. | No accuracy gain; frequently reveals that the model is using a shortcut |
| Alternate backbones | Repeat the full pipeline with EfficientNet-B3 and MobileNetV3-Large. Tabulate accuracy, parameters, and inference latency. | Turns "which backbone" from a guess into a measurement |
The Grad-CAM extension is the most instructive and the least glamorous. Fine-grained flower datasets are full of background correlations — one species photographed mostly in gardens, another mostly against sky — and a model can score well by learning the background. If Grad-CAM shows the heat concentrated outside the flower, your test accuracy will not survive contact with a new photographer.
The mistakes this project is designed to catch
| Pitfall | Symptom | Fix |
|---|---|---|
| Using the larger split as training data | Suspiciously high accuracy, incomparable to published numbers | Print all three split sizes and respect the dataset's design |
| Replacing the head before freezing | Trainable parameters read 100%; overfits immediately | Freeze first, replace second |
| Augmenting the validation set | Validation accuracy noisy and below training by an odd amount | Separate transform objects for train and eval |
| From-scratch learning rate on pretrained weights | Loss spikes in the first 50 steps, then flattens at chance | 1e-4 or lower, and train the head first |
| Wrong preprocessing constants | Frozen baseline several points below expectation | Read them from weights.transforms() |
| Tuning on the test set | Test accuracy stops predicting real-world performance | Tune on validation; touch test once |
| Keeping the last epoch instead of the best | Final model worse than a number you saw mid-training | Checkpoint on best validation and reload it |
| Reporting accuracy alone on imbalanced data | Respectable headline, rare classes never predicted | Report macro-F1 and per-class recall |
What to carry out of this into real work
The thing worth internalising is not the ResNet recipe — backbones change every eighteen months. It is the discipline of the staged pipeline, which transfers to every model you will ever fine-tune.
- Establish a cheap, hard-to-break baseline first. The frozen backbone takes twelve minutes and produces a number that every subsequent idea must beat. Without it, you have no way to tell whether your clever schedule helped or quietly hurt.
- Change one thing per phase and validate between phases. When phase 3 beats phase 2 you know exactly why. Change three things at once and you have learned nothing, whichever way the number moves.
- Treat a failed phase as a measurement. If unfreezing more layers makes validation worse, the dataset has told you its capacity limit. That is a finding about your data, and the response is more data or better augmentation, not a longer hyperparameter sweep.
- Spend your last hour on errors, not on accuracy. The confused-pair analysis above turned an 80.8% headline into a specific hypothesis about colour saturation. Findings like that come from looking at images; they never come from looking at a scalar.
- Report honestly. One test-set evaluation, at the end, with the split sizes stated. A number that is 4 points lower and trustworthy is worth more than one that is 4 points higher and unreproducible — because the trustworthy one is the one that predicts what happens after you ship.