Transfer learning
Be able to reuse a pretrained model, freeze layers and fine-tune the head.
Prerequisites
- DTrain a neural network in PyTorchrequired
Intuition
Training an image model from scratch takes hundreds of thousands of images. But a model already trained on ImageNet has learnt edges, textures and shapes — things that hold for all images. Just swap the last layer and teach it your classes.
Three strategies:
| Strategy | When | How |
|---|---|---|
| Feature extraction | very little data (< 1 000) | freeze everything, train only a new head |
| Fine-tune the top | moderate (1 000–10 000) | freeze the early layers, train the last ones plus the head |
| Full fine-tuning | a lot of data, or a different domain | train everything at a low lr |
The more your data resembles the pretraining data, the more you can freeze.
Code
import torch, torch.nn as nn
from torchvision import models
m = models.resnet18(weights="IMAGENET1K_V1")
for p in m.parameters():
p.requires_grad = False # freeze everything
m.fc = nn.Linear(m.fc.in_features, 5) # a new head, 5 classes (trained)
opt = torch.optim.AdamW(m.fc.parameters(), lr=1e-3)
# … train the head for a few epochs …
# Step 2: unfreeze the last blocks and fine-tune with a much lower lr
for p in m.layer4.parameters():
p.requires_grad = True
opt = torch.optim.AdamW([
{"params": m.layer4.parameters(), "lr": 1e-5},
{"params": m.fc.parameters(), "lr": 1e-4},
])
Three mistakes that cost percentage points:
- The wrong normalisation. Use the same mean and std as the pretraining (ImageNet: mean
[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]). - Too high an lr on the pretrained layers. 1e-5 to 1e-4, otherwise what the model knows is erased.
- Forgetting
model.eval()— batchnorm in a resnet updates its statistics during training even for frozen layers unless you handle it.
Mastery means
- Reuses a pretrained model
- Chooses between freezing and full fine-tuning
- Avoids the common mistakes with normalisation and the lr
Sign in to do the exercises and build your mastery up.
Sources
- PyTorch — Transfer learning tutorial (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0