Course Content
Transfer Learning and Pretraining
3 sections · 7 lessons
Feature Extraction vs Fine-Tuning — What to Freeze
An engineer has 420 chest X-rays labelled pneumonia or normal. She loads a ResNet-50 with ImageNet weights, swaps the final layer for a two-class output, and calls model.parameters() in her optimiser with the learning rate she always uses: lr=1e-3, SGD with momentum. She hits train.
Step 1: loss 0.71. Step 12: loss 4.83. Step 40: loss 0.69 and completely flat. Final validation accuracy: 58% — below the 62% she would get by predicting "normal" for every image, which is to say worse than useless.
Her colleague runs a different version. Same model, same weights, same data. He freezes every layer except the final one, trains for eight epochs, and gets 84% validation accuracy in under three minutes.
Same pretrained model. Same dataset. A 26-point difference, and the only distinction is which weights were allowed to move. That single choice — freeze or update, and if update then how much and how fast — is the most consequential decision in transfer learning, and it has two named answers.
Feature extraction: treat the network as a fixed measuring instrument
Feature extraction means you freeze the pretrained backbone entirely and train only a new classifier on top of its outputs. The backbone becomes a fixed function that turns an image into a vector of numbers. You never change it. You only learn how to read it.
Concretely, for ResNet-50: strip off the 1000-way ImageNet classifier, keep everything up to and including global average pooling, and you have a function that maps a 224×224×3 image to a 2048-dimensional vector. Those 2048 numbers describe the image in terms the network learned from 1.28 million photographs — presence of textures, parts, shapes, materials. Your job is to learn a mapping from that vector to your labels.
The parameter count shows why this is safe on small data. Full ResNet-50 is 25.6 million parameters; a linear head from 2048 features to 2 classes is 2048×2+2=4,098. With 420 images that is ten parameters per example rather than sixty thousand.
Feature extraction converts an impossible optimisation problem — fit 25.6 million weights from 420 examples — into an easy one: fit a logistic regression on 2048 well-chosen features.
Doing it in PyTorch
1import torch2import torch.nn as nn3from torchvision import models4from torchvision.models import ResNet50_Weights56# Load pretrained weights. The `weights=` enum is the current API;7# the old `pretrained=True` flag is deprecated and prints a warning.8weights = ResNet50_Weights.IMAGENET1K_V29model = models.resnet50(weights=weights)1011# 1. Freeze every parameter in the network.12for param in model.parameters():13 param.requires_grad = False1415# 2. Replace the head. Newly created layers have requires_grad=True16# by default, so this un-freezes exactly what we want.17num_features = model.fc.in_features # 2048 for ResNet-5018model.fc = nn.Linear(num_features, 2)1920# 3. Give the optimiser ONLY the trainable parameters. Passing21# model.parameters() also works (parameters with no gradient are22# skipped), but this makes the intent explicit and easy to count.23trainable = [p for p in model.parameters() if p.requires_grad]24print(f"trainable: {sum(p.numel() for p in trainable):,}") # 4,0982526optimiser = torch.optim.AdamW(trainable, lr=1e-3, weight_decay=1e-4)27criterion = nn.CrossEntropyLoss()Two traps live in that snippet. First, the order matters: freeze first, then replace the head. Do it the other way round and the loop sets requires_grad=False on your brand-new head too, and you will train nothing at all — the loss will sit perfectly still and you will spend an hour wondering why.
Second, if the model contains BatchNorm layers — ResNet does, in abundance — requires_grad=False does not stop BatchNorm from updating its running mean and variance. Those are buffers, not parameters. In training mode they keep absorbing your target-domain statistics, which silently changes the "frozen" backbone's behaviour. If you want a genuinely fixed feature extractor, put the backbone in eval mode:
1model.train() # sets everything to training mode2# ...then force the frozen backbone back to eval so BatchNorm3# uses its ImageNet running statistics and stops updating them.4for module in model.modules():5 if isinstance(module, nn.BatchNorm2d):6 module.eval()Which behaviour you want is a genuine judgement call — letting BatchNorm re-estimate statistics on your data is itself a mild form of domain adaptation and often helps. The failure is not knowing which one you got.
Doing it in Keras
1import tensorflow as tf2from tensorflow.keras import layers, Model34base = tf.keras.applications.ResNet50(5 weights="imagenet",6 include_top=False, # drop the 1000-way classifier7 input_shape=(224, 224, 3),8)9base.trainable = False # freezes weights AND puts BN in inference mode1011inputs = tf.keras.Input(shape=(224, 224, 3))12x = tf.keras.applications.resnet50.preprocess_input(inputs)13x = base(x, training=False) # belt-and-braces: keep BN in inference mode14x = layers.GlobalAveragePooling2D()(x)15x = layers.Dropout(0.3)(x)16outputs = layers.Dense(2, activation="softmax")(x)1718model = Model(inputs, outputs)19model.compile(optimizer=tf.keras.optimizers.Adam(1e-3),20 loss="sparse_categorical_crossentropy",21 metrics=["accuracy"])Keras is stricter here in a helpful way: setting base.trainable = False also forces BatchNorm into inference mode, so the Keras version of "frozen" really is frozen.
The trick that makes feature extraction almost free
If the backbone never changes, then for a given image its 2048-dimensional feature vector never changes either. So compute it once, cache it, and train the head on the cached vectors. Your training loop stops touching the GPU-heavy convolutional stack entirely.
1import numpy as np23backbone = nn.Sequential(*list(model.children())[:-1]).eval().cuda()45feats, labels = [], []6with torch.no_grad():7 for images, y in loader: # ONE pass over the data8 f = backbone(images.cuda()).flatten(1) # (B, 2048)9 feats.append(f.cpu()); labels.append(y)1011X = torch.cat(feats).numpy() # (420, 2048)12y = torch.cat(labels).numpy()1314from sklearn.linear_model import LogisticRegression15clf = LogisticRegression(max_iter=2000, C=1.0).fit(X, y)The arithmetic on 420 images: one forward pass costs about 4 GFLOPs, so 1.7 TFLOPs total — under two seconds on a modern GPU. Fitting the logistic regression on a 420×2048 matrix takes under a second on CPU. Fine-tuning instead runs a forward and backward pass over all 420 images, 20 times over: roughly 60× the compute.
The catch: caching kills random augmentation, because augmentation must happen before the backbone and you have moved the backbone out of the loop. At 420 images augmentation matters, so either cache several augmented copies of each image or accept the online cost.
What you gain and what you give up
| Advantages of feature extraction | Disadvantages |
|---|---|
| Very fast — minutes, sometimes seconds with cached features | Accuracy ceiling is lower when the target domain differs from the source |
| Almost impossible to overfit with a linear head on 2048 features | Cannot adapt features to your domain at all — X-ray textures stay described in ImageNet terms |
| Tiny memory footprint; trains fine on CPU or a laptop GPU | Sensitive to input preprocessing mismatches, with no ability to compensate |
| Highly reproducible — very few moving parts | Wastes capacity when you do have enough data to fine-tune |
| Gives a trustworthy baseline that fine-tuning must beat | Frozen BatchNorm statistics can be badly calibrated for your data |
Fine-tuning: let the features move
Fine-tuning means allowing some or all of the pretrained weights to update during training on your task. You are no longer just learning to read the features; you are reshaping them.
That is powerful, and it is exactly what went wrong for the engineer in the opening. Fine-tuning done carelessly is worse than not fine-tuning at all, and the mechanism is precise enough to be worth stating exactly.
Why a high learning rate destroys the model in the first fifty steps
Your new head is randomly initialised. At step 1 it produces essentially random logits, so the cross-entropy loss is around ln(2)≈0.69 for two balanced classes — and if the random init happens to be confidently wrong, considerably higher. Those large errors backpropagate through the whole backbone.
The pretrained weights encode features that took roughly 115 million image presentations to build. A single SGD step at lr=1e-3 with gradients driven by a random head can move a weight by more than the entire range it explored during the last twenty epochs of pretraining. Do that forty times and the features are gone. The loss then settles at 0.69 — the model has found the trivial solution of predicting the class prior — and it never recovers, because now it genuinely is training from scratch on 420 images.
Catastrophic forgetting is not a slow drift. It happens in the first dozen steps, driven by the random head, and by the time you look at a validation number the pretrained features have already been overwritten.
Two fixes address it directly. First, warm up the head: train with the backbone frozen for one or two epochs so the head becomes sensible, then unfreeze. Now the gradients entering the backbone are small and meaningful. Second, use a much lower learning rate for pretrained weights — typically 10× to 100× lower than you would use from scratch.
Full fine-tuning
1for param in model.parameters():2 param.requires_grad = True34optimiser = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)Every weight moves. This has the highest ceiling and the highest risk. It is the right choice when you have thousands of examples per class and your domain differs meaningfully from the source. On 420 X-rays it will overfit unless you regularise hard.
Selective fine-tuning by progressive unfreezing
The better default is to unfreeze gradually, from the top down. The reasoning follows directly from how the feature hierarchy is organised: early layers detect edges and textures, which are as valid for X-rays as for photographs and therefore need no adjustment. Late layers encode whole-object ImageNet concepts, which are useless to you and need the most change.
1def freeze_all_but(model, groups_to_train):2 """groups_to_train: list of nn.Module attributes to unfreeze."""3 for p in model.parameters():4 p.requires_grad = False5 for g in groups_to_train:6 for p in g.parameters():7 p.requires_grad = True89schedule = [10 # (epochs, unfrozen groups, lr)11 (3, [model.fc], 1e-3),12 (4, [model.fc, model.layer4], 1e-4),13 (4, [model.fc, model.layer4, model.layer3], 5e-5),14]1516for n_epochs, groups, lr in schedule:17 freeze_all_but(model, groups)18 opt = torch.optim.AdamW(19 [p for p in model.parameters() if p.requires_grad],20 lr=lr, weight_decay=1e-4)21 for _ in range(n_epochs):22 train_one_epoch(model, loader, opt, criterion)23 validate(model, val_loader) # check before widening furtherThe validation check between phases is the whole point. If unfreezing layer3 makes validation accuracy fall, stop — you have found the depth at which your dataset can no longer support more free parameters.
The unfreezing pyramid
How deep you should go is mostly a function of dataset size. This table is a starting point, not a law, but it is a good starting point.
| Examples per class | What to unfreeze | Approx. trainable params (ResNet-50) | Typical LR |
|---|---|---|---|
| Under 100 | Head only | ~4K–20K | 1e-3 |
| 100–500 | Head + layer4 | ~15M | 1e-4 |
| 500–2,000 | Head + layer4 + layer3 | ~22M | 5e-5 to 1e-4 |
| 2,000–10,000 | Everything except conv1 and layer1 | ~25M | 3e-5 to 5e-5 |
| Over 10,000 | Everything | 25.6M | 1e-5 to 3e-5, with warmup |
Note that layer4 alone is about 15 million of ResNet-50's 25.6 million parameters — the network is extremely top-heavy. Unfreezing just the last block already puts most of the model in play, which is why the jump from "head only" to "head + layer4" is where overfitting usually first appears.
Discriminative learning rates
Progressive unfreezing is a coarse instrument: a layer is either fully trainable or fully frozen. Discriminative learning rates are the smooth version — every layer trains, but deeper layers get exponentially smaller learning rates.
Pick a base rate for the head and divide by a constant factor per group going down. A factor of 2.6 per group is a common choice, popularised by the ULMFiT work; factors between 2 and 10 all behave sensibly.
1base_lr = 1e-32factor = 2.634groups = [model.fc, model.layer4, model.layer3, model.layer2, model.layer1]5param_groups = []6for depth, group in enumerate(groups):7 param_groups.append({8 "params": group.parameters(),9 "lr": base_lr / (factor ** depth),10 })1112optimiser = torch.optim.AdamW(param_groups, weight_decay=1e-4)The resulting rates, computed exactly:
| Group | Formula | Learning rate | Relative to head |
|---|---|---|---|
fc (new head) | 10−3/2.60 | 1.00e-3 | 1× |
layer4 | 10−3/2.61 | 3.85e-4 | 0.38× |
layer3 | 10−3/2.62 | 1.48e-4 | 0.15× |
layer2 | 10−3/2.63 | 5.69e-5 | 0.057× |
layer1 | 10−3/2.64 | 2.19e-5 | 0.022× |
The bottom of the network moves about 46 times more slowly than the top. Edge detectors are nudged; the classifier is rebuilt. This usually outperforms both hard freezing and uniform fine-tuning, because it removes the artificial cliff at the freeze boundary — and that cliff is exactly where co-adapted neighbouring layers get cut apart from each other.
Choosing between them
The theoretical comparison
| Dimension | Feature extraction | Fine-tuning |
|---|---|---|
| Trainable parameters | Thousands | Millions to tens of millions |
| Data required | Tens per class is workable | Hundreds to thousands per class |
| Training time (420 images) | Seconds to a few minutes | Tens of minutes to hours |
| GPU memory | Low — no backbone gradients or optimiser state | High — gradients plus optimiser state for every weight |
| Overfitting risk | Very low | High without regularisation |
| Accuracy ceiling | Limited by how well source features describe your data | Substantially higher when data supports it |
| Sensitivity to learning rate | Forgiving | Unforgiving — the single biggest cause of failure |
| Handles large domain shift | Poorly | Well, given enough data |
| Reproducibility | Excellent | Noticeable run-to-run variance |
The empirical comparison
Here is the same binary cats-versus-dogs problem — a task very close to ImageNet, since ImageNet contains 120 dog breeds and several cat breeds — run at three dataset sizes with a ResNet-50. The figures are illustrative, typical of what this experiment produces rather than a published benchmark; your own numbers will differ by a point or two, but the pattern holds.
| Training images | From scratch | Feature extraction | Fine-tune layer4+head | Full fine-tune |
|---|---|---|---|---|
| 200 | 56.5% | 96.8% | 96.1% | 91.2% |
| 2,000 | 72.3% | 97.4% | 98.6% | 98.1% |
| 20,000 | 91.8% | 97.9% | 99.0% | 99.3% |
Three things to read out of that table. At 200 images, full fine-tuning is worse than doing nothing to the backbone — 5.6 points worse — because 25.6 million parameters cannot be constrained by 200 examples. At 2,000 the optimum has moved to partial fine-tuning. At 20,000 full fine-tuning finally wins, and from-scratch training has closed most but not all of the gap.
Note also how flat the feature-extraction column is: 96.8% to 97.9% across a hundredfold increase in data. Frozen features cannot exploit more data — once the linear head converges there is nothing left to learn. That flatness is the signature of a capacity limit.
The decision procedure
How similar is your data to the pretraining source? (natural photos, standard objects, RGB, ~224px) SIMILAR DIFFERENT ------- ---------SMALL Feature extraction. Feature extraction, butDATASET Frozen features already from an EARLIER layer -(< ~1k) describe your data well. late layers are too Full FT will overfit. source-specific to help. Consider a closer source.LARGE Fine-tune the top few Fine-tune deeply, or allDATASET blocks. Diminishing of it. This is where the(> ~10k) returns from going deeper. domain gap gets closed.The top-right cell is the hard one, and the honest answer is that no freezing strategy rescues it — small dataset plus distant domain needs a better-matched pretrained source, or self-supervised pretraining on your own unlabelled data.
Practices that consistently pay off
Always establish the frozen baseline first
Ten minutes of feature extraction gives you a number every later experiment must beat. Without it you cannot tell whether a fine-tuning schedule is helping or quietly hurting. The engineer above spent two days on the 58% model; the three-minute run would have flagged the problem immediately.
Warm up the head before unfreezing anything
One or two frozen epochs cost almost nothing and prevent the entire catastrophic-forgetting failure mode. Treat it as mandatory.
Use learning rates 10–100× below from-scratch values
| Setting | Typical LR (SGD) | Typical LR (AdamW) |
|---|---|---|
| Training ResNet from scratch | 0.1 | 1e-3 |
| New head only, backbone frozen | 1e-2 | 1e-3 |
| Fine-tuning the top block | 1e-3 | 1e-4 |
| Full fine-tuning | 1e-4 | 1e-5 to 3e-5 |
Monitor the train–validation gap, not just validation accuracy
Validation accuracy alone tells you where you are; the gap tells you where you are heading. A gap under about 5 points means you have headroom and can safely unfreeze more. A gap over 15 points means you are already memorising and unfreezing more will make it worse. Log both every epoch.
Prefer discriminative rates over hard freezing
Hard freezing creates a discontinuity — a boundary across which co-adapted features cannot adjust to each other. Discriminative rates achieve the same protection smoothly, and typically land half a point to two points higher.
Keeping fine-tuning from overfitting
Once weights are unfrozen, the regularisation you apply decides whether it works.
Augmentation, matched to the domain
1from torchvision import transforms23train_tf = transforms.Compose([4 transforms.RandomResizedCrop(224, scale=(0.7, 1.0)),5 transforms.RandomHorizontalFlip(),6 transforms.RandomRotation(15),7 transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),8 transforms.ToTensor(),9 transforms.Normalize([0.485, 0.456, 0.406],10 [0.229, 0.224, 0.225]),11 transforms.RandomErasing(p=0.25, scale=(0.02, 0.15)),12])Two domain-specific warnings. RandomHorizontalFlip is free on cats and dogs and harmful on chest X-rays, where mirroring moves the heart to the wrong side and teaches the model that dextrocardia is normal. ColorJitter is meaningless on greyscale radiographs and destructive on histopathology, where stain colour carries diagnostic signal. Every transform must preserve the label.
Also note that the normalisation constants above are ImageNet's channel means and standard deviations. If your fine-tuning preprocessing does not match the preprocessing used during pretraining, the frozen features are being fed inputs from a distribution they have never seen, and you can lose several points before you have done anything else wrong.
Regularisation in the head, and weight decay in the backbone
1model.fc = nn.Sequential(2 nn.Dropout(0.5),3 nn.Linear(2048, 512),4 nn.BatchNorm1d(512),5 nn.ReLU(inplace=True),6 nn.Dropout(0.3),7 nn.Linear(512, num_classes),8)Higher dropout on the layer closest to the features; lower on the smaller one. Weight decay of 1e-4 to 1e-2 pulls weights towards zero; a stronger variant penalises distance from the pretrained weights instead, regularising towards the source model rather than towards nothing.
Train fewer epochs than you think
Fine-tuning converges fast. On a few hundred images, 10–20 epochs is usually the whole budget; the best validation score often arrives at epoch 6 or 7. Use early stopping with a patience of 3–5 epochs, and always restore the best checkpoint rather than the last one.
1best_acc, patience, bad_epochs = 0.0, 5, 02for epoch in range(30):3 train_one_epoch(model, train_loader, optimiser, criterion)4 acc = validate(model, val_loader)5 if acc > best_acc:6 best_acc, bad_epochs = acc, 07 torch.save(model.state_dict(), "best.pt")8 else:9 bad_epochs += 110 if bad_epochs >= patience:11 break12model.load_state_dict(torch.load("best.pt"))The whole thing, end to end
Here is the complete recipe on the engineer's 420 X-rays, written the way it should have been written the first time.
1import torch, torch.nn as nn2from torchvision import models3from torchvision.models import ResNet50_Weights45model = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)6model.fc = nn.Sequential(nn.Dropout(0.4), nn.Linear(2048, 2))7model = model.cuda()8criterion = nn.CrossEntropyLoss(label_smoothing=0.05)910def set_trainable(groups):11 for p in model.parameters():12 p.requires_grad = False13 for g in groups:14 for p in g.parameters():15 p.requires_grad = True1617def make_opt(lr):18 return torch.optim.AdamW(19 [p for p in model.parameters() if p.requires_grad],20 lr=lr, weight_decay=1e-4)2122# Phase 1 - head only. Establishes the baseline AND warms up the head23# so that phase 2 does not blow up the pretrained features.24set_trainable([model.fc])25run(epochs=5, opt=make_opt(1e-3)) # expect ~84%2627# Phase 2 - add the last residual block at a 10x lower rate.28set_trainable([model.fc, model.layer4])29run(epochs=8, opt=make_opt(1e-4)) # expect ~89%3031# Phase 3 - discriminative rates across the whole network.32for p in model.parameters():33 p.requires_grad = True34opt = torch.optim.AdamW([35 {"params": model.fc.parameters(), "lr": 1e-4},36 {"params": model.layer4.parameters(), "lr": 3.8e-5},37 {"params": model.layer3.parameters(), "lr": 1.5e-5},38 {"params": model.layer2.parameters(), "lr": 5.7e-6},39 {"params": model.layer1.parameters(), "lr": 2.2e-6},40], weight_decay=1e-4)41run(epochs=10, opt=opt) # expect ~91%Each phase must be validated before the next begins. If phase 2 does not beat phase 1, do not run phase 3 — your dataset has told you it cannot support that many free parameters, and the correct response is to keep the phase 1 model and spend the effort on augmentation or on collecting data.
What this means when you build something
The mistake almost everyone makes at the start is treating "fine-tuning" as the sophisticated option and feature extraction as the beginner's shortcut. The data says the opposite: below roughly a thousand examples, frozen features beat full fine-tuning outright, and the gap widens as the dataset shrinks. Sophistication is choosing correctly, not choosing the expensive thing.
So the operational rule is a ladder, and you climb it one rung at a time with a validation number at each step:
- Frozen backbone, linear head. Ten minutes. This is your floor and your sanity check. If this is near chance, something is broken in your data pipeline — wrong normalisation, shuffled labels, corrupt images — and no amount of fine-tuning will save you.
- Frozen backbone, small MLP head with dropout. Occasionally worth a point or two when class boundaries are not linearly separable in feature space.
- Unfreeze the last block, learning rate 10× lower. The single highest-value step for most projects.
- Discriminative learning rates across the full network. Worth doing once the previous rung has clearly helped.
Stop climbing the moment a rung fails to improve validation accuracy. That failure is information: it tells you that you have reached the capacity your dataset can support, and further effort belongs in the data — more examples, better augmentation, cleaner labels — rather than in the optimiser. The engineer's original run skipped straight to rung four with a from-scratch learning rate, which is why it produced a model that had forgotten ImageNet and learned nothing to replace it.