Transfer Learning and Pretraining

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.

Ten images per class, run honestly end to endSplit first,and check forduplicatesAugment inproportion to10 per classFrozen baselineyou have to beatUnfreeze instages,rates per blockConfusionmatrix, thenread the errors98.9% on this dataset usually means the split leaked, not that the model is better.
The frozen baseline is the control: without it, every later number is an improvement over nothing in particular.

The project at a glance

StageWhat you buildTimeTarget result
1. Data preparationLoaders, split verification, augmentation sized to the data~1 hourCorrect splits, verified visually
2. Model constructionPretrained backbone, replaced head, reusable train/eval loop~1 hourForward pass with the right output shape
3. Frozen baselineFeature extraction, backbone frozen~2 hoursRoughly 80–87% test accuracy
4. Progressive fine-tuningStaged unfreezing, discriminative learning rates~2 hoursRoughly 90–95% test accuracy
5. Evaluation and analysisConfusion matrix, per-class metrics, error inspection~1 hourA written account of what the model gets wrong

Dataset options

DatasetClassesTrain imagesDifficultyWhy choose it
Oxford Flowers-1021021,020HardGenuinely small data — 10 per class. The recommended default.
Oxford-IIIT Pet373,680MediumBreed-level fine-grained, close to ImageNet, forgiving
Stanford Cars1968,144HardVery fine-grained; model year matters more than shape
Food-10110175,750MediumLarge enough that full fine-tuning genuinely wins
Caltech-101101~3,000EasyFastest 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.

Text
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.md

Part 1: data preparation

Loading the data, and checking the split

Python
from torchvision import datasetsfrom torchvision.models import ResNet50_Weightsweights = ResNet50_Weights.IMAGENET1K_V2train_raw = datasets.Flowers102(root="data", split="train", download=True)val_raw   = datasets.Flowers102(root="data", split="val",   download=True)test_raw  = datasets.Flowers102(root="data", split="test",  download=True)print(len(train_raw), len(val_raw), len(test_raw))   # 1020 1020 6149

Print 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

Python
from collections import Counterimport numpy as npcounts = Counter(label for _, label in train_raw)c = np.array(sorted(counts.values()))print(f"classes {len(counts)}  min {c.min()}  max {c.max()}  median {np.median(c)}")sizes = [img.size for img, _ in list(train_raw)[:200]]ws, hs = zip(*sizes)print(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

Python
from torchvision import transformsnorm = transforms.Normalize(mean=[0.485, 0.456, 0.406],                            std=[0.229, 0.224, 0.225])train_tf = transforms.Compose([    transforms.RandomResizedCrop(224, scale=(0.5, 1.0), ratio=(0.75, 1.33)),    transforms.RandomHorizontalFlip(),    transforms.RandomVerticalFlip(),          # flowers have no fixed "up"    transforms.RandomRotation(30),    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.05),    transforms.ToTensor(), norm,    transforms.RandomErasing(p=0.25, scale=(0.02, 0.2)),])eval_tf = transforms.Compose([    transforms.Resize(232),                   # matches the V2 weights recipe    transforms.CenterCrop(224),    transforms.ToTensor(), norm,])

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.

Python
from torch.utils.data import Dataset, DataLoaderclass TransformedDataset(Dataset):    def __init__(self, base, transform):        self.base, self.transform = base, transform    def __len__(self):        return len(self.base)    def __getitem__(self, i):        img, label = self.base[i]        return self.transform(img), labeltrain_ds = TransformedDataset(train_raw, train_tf)val_ds   = TransformedDataset(val_raw,   eval_tf)test_ds  = TransformedDataset(test_raw,  eval_tf)train_dl = DataLoader(train_ds, batch_size=32, shuffle=True,                      num_workers=4, pin_memory=True, drop_last=True)val_dl   = DataLoader(val_ds,  batch_size=64, num_workers=4, pin_memory=True)test_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

Python
import torch, torch.nn as nnfrom torchvision import modelsdef build_model(num_classes=102, dropout=0.4):    weights = models.ResNet50_Weights.IMAGENET1K_V2    model = models.resnet50(weights=weights)    for p in model.parameters():        # freeze FIRST        p.requires_grad = False    model.fc = nn.Sequential(           # then replace - new layers train        nn.Dropout(dropout),        nn.Linear(model.fc.in_features, num_classes),    )    return modelmodel = build_model().cuda()total = sum(p.numel() for p in model.parameters())trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)print(f"{trainable:,} / {total:,} trainable ({100*trainable/total:.3f}%)")# 208,998 / 23,717,030 trainable (0.881%)with torch.no_grad():    print(model(torch.randn(2, 3, 224, 224).cuda()).shape)   # [2, 102]

Verify both printed lines. The head is 2048×102+102=208,9982048 \times 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

Python
def run_epoch(model, loader, criterion, optimiser=None, device="cuda"):    training = optimiser is not None    model.train(training)    total_loss, correct, n = 0.0, 0, 0    for x, y in loader:        x, y = x.to(device, non_blocking=True), y.to(device)        with torch.set_grad_enabled(training):            out = model(x)            loss = criterion(out, y)        if training:            optimiser.zero_grad(set_to_none=True)            loss.backward()            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)            optimiser.step()        total_loss += loss.item() * y.size(0)        correct += (out.argmax(1) == y).sum().item()        n += y.size(0)    return total_loss / n, correct / ndef train_phase(model, name, optimiser, epochs, criterion,                train_dl, val_dl, patience=8):    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimiser, T_max=epochs)    best, bad = 0.0, 0    for epoch in range(epochs):        tr_loss, tr_acc = run_epoch(model, train_dl, criterion, optimiser)        va_loss, va_acc = run_epoch(model, val_dl, criterion)        scheduler.step()        print(f"[{name}] {epoch:2d}  train {tr_acc:.4f}  val {va_acc:.4f}  "              f"gap {tr_acc - va_acc:+.4f}")        if va_acc > best:            best, bad = va_acc, 0            torch.save(model.state_dict(), f"checkpoints/{name}.pt")        else:            bad += 1            if bad >= patience:                break    model.load_state_dict(torch.load(f"checkpoints/{name}.pt"))    return best

Print 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

Python
criterion = nn.CrossEntropyLoss(label_smoothing=0.1)optimiser = torch.optim.AdamW(    [p for p in model.parameters() if p.requires_grad],    lr=1e-3, weight_decay=1e-4)best_frozen = train_phase(model, "phase1_frozen", optimiser,                          epochs=30, criterion=criterion,                          train_dl=train_dl, val_dl=val_dl)print(f"frozen baseline: {best_frozen:.4f}")     # validation; expect mid-0.80s

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

Python
def set_trainable(model, groups):    for p in model.parameters():        p.requires_grad = False    for g in groups:        for p in g.parameters():            p.requires_grad = True# Phase 2: add layer4 only, at a 10x lower learning rate.set_trainable(model, [model.fc, model.layer4])opt2 = torch.optim.AdamW(    [p for p in model.parameters() if p.requires_grad],    lr=1e-4, weight_decay=1e-4)best_p2 = train_phase(model, "phase2_layer4", opt2, epochs=25,                      criterion=criterion, train_dl=train_dl, val_dl=val_dl)print(f"phase 2: {best_p2:.4f}")      # expect ~0.92

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

Python
for p in model.parameters():    p.requires_grad = Truebase_lr, factor = 1e-4, 2.6groups = [model.fc, model.layer4, model.layer3, model.layer2, model.layer1]opt3 = torch.optim.AdamW(    [{"params": g.parameters(), "lr": base_lr / factor ** d}     for d, g in enumerate(groups)],    weight_decay=1e-4)best_p3 = train_phase(model, "phase3_discriminative", opt3, epochs=25,                      criterion=criterion, train_dl=train_dl, val_dl=val_dl)print(f"phase 3: {best_p3:.4f}")      # expect ~0.94
GroupFormulaLearning rateRelative to head
fc10−4/2.6010^{-4} / 2.6^{0}1.00e-41×
layer410−4/2.6110^{-4} / 2.6^{1}3.85e-50.38×
layer310−4/2.6210^{-4} / 2.6^{2}1.48e-50.15×
layer210−4/2.6310^{-4} / 2.6^{3}5.69e-60.057×
layer110−4/2.6410^{-4} / 2.6^{4}2.19e-60.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.

Python
import numpy as npfrom sklearn.metrics import confusion_matrix, classification_report@torch.no_grad()def predict(model, loader, device="cuda"):    model.eval()    preds, targets, confs = [], [], []    for x, y in loader:        probs = model(x.to(device)).softmax(1)        conf, pred = probs.max(1)        preds.append(pred.cpu()); targets.append(y); confs.append(conf.cpu())    return (torch.cat(preds).numpy(), torch.cat(targets).numpy(),            torch.cat(confs).numpy())preds, targets, confs = predict(model, test_dl)cm = confusion_matrix(targets, preds)print(classification_report(targets, preds, digits=3))# The most confused pairs, ignoring the diagonal.off = cm.copy(); np.fill_diagonal(off, 0)for _ in range(10):    i, j = np.unravel_index(off.argmax(), off.shape)    print(f"true {i:3d} -> predicted {j:3d} : {off[i, j]} times")    off[i, j] = 0

Reading 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 →SunflowerDaisyConeflowerBlack-eyed SusanTotal
Sunflower5814265
Daisy2601265
Coneflower50471365
Black-eyed Susan31164565

Overall accuracy is (58+60+47+45)/260=210/260=80.8%(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=684 + 1 + 47 + 16 = 68 predictions, of which 47 are correct, so precision is 47/68=69.1%47/68 = 69.1\%. Read across the Coneflower row: 65 true examples, 47 found, so recall is 47/65=72.3%47/65 = 72.3\%, giving an F1 of 70.7%. Compare that with Daisy, whose precision is 60/62=96.8%60/62 = 96.8\% and recall 60/65=92.3%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

Python
wrong = np.where(preds != targets)[0]order = wrong[np.argsort(-confs[wrong])]        # most confident mistakes firstfor i in order[:20]:    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

DeliverableWhat it must contain
Working codeRuns end to end from a clean checkout; fixed random seeds; the split sizes printed and unaltered
Training historyPer-epoch train and validation accuracy for all three phases, in one plot with phase boundaries marked
Results tableTest accuracy and macro-F1 for each phase, so the gain from each stage is visible
Confusion analysisThe ten most confused class pairs, with images, and a hypothesis for each
Error inspectionTwenty highest-confidence errors, categorised as label noise / ambiguous / model failure
Written reportWhat you tried, what failed and why, what you would do with another week
CriterionWeightWhat full marks looks like
Correctness of the pipeline25%Splits respected, no leakage, augmentation on training only, preprocessing matched to the weights
Staged methodology25%Frozen baseline established first; each phase validated before the next; losing phases discarded
Final performance20%Test accuracy above the frozen baseline by a clear margin
Error analysis20%Specific, image-level findings with plausible causes — not a restated accuracy
Reproducibility10%Seeds fixed, environment pinned, results reproduced within a point on a second run

Optional extensions

ExtensionWhat to doRealistic gain
EnsemblingAverage 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 distillationTrain 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
DeploymentExport 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
ExplanationRun 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 backbonesRepeat 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

PitfallSymptomFix
Using the larger split as training dataSuspiciously high accuracy, incomparable to published numbersPrint all three split sizes and respect the dataset's design
Replacing the head before freezingTrainable parameters read 100%; overfits immediatelyFreeze first, replace second
Augmenting the validation setValidation accuracy noisy and below training by an odd amountSeparate transform objects for train and eval
From-scratch learning rate on pretrained weightsLoss spikes in the first 50 steps, then flattens at chance1e-4 or lower, and train the head first
Wrong preprocessing constantsFrozen baseline several points below expectationRead them from weights.transforms()
Tuning on the test setTest accuracy stops predicting real-world performanceTune on validation; touch test once
Keeping the last epoch instead of the bestFinal model worse than a number you saw mid-trainingCheckpoint on best validation and reload it
Reporting accuracy alone on imbalanced dataRespectable headline, rare classes never predictedReport 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.

  1. 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.
  2. 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.
  3. 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.
  4. 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.
  5. 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.