Training a modern image or text model from random weights needs far more labelled data and compute than most teams have. Transfer learning sidesteps that: start from a network trained on a large source task, keep what it learned about the structure of images or language, and adapt it to your smaller target task. Done well, a few thousand labelled examples go a long way. Done badly, the pretrained knowledge is destroyed in the first few hundred steps, or a mismatch between source and target makes the result worse than a small model trained from scratch.
This article covers the classical setting: a pretrained convolutional network or encoder adapted to a new classification or regression task. It builds from why features transfer to a working two-stage PyTorch recipe, then covers the details that decide whether it works: BatchNorm, learning rates, preprocessing and evaluation. If you are adapting a large language model, the same ideas apply but the tooling is different; start with parameter-efficient fine-tuning after reading this.
Why features transfer at all
A deep network trained on a large, varied dataset learns a hierarchy. Yosinski and colleagues showed in 2014, by transplanting layers between networks trained on different halves of ImageNet, that the first layers learn features close to universal: oriented edges, colour blobs and textures. Later layers become increasingly specific to the source classes. Transferability drops as you go up the stack, and it drops faster when the source and target tasks are further apart.
That gives the core intuition for every decision that follows. The lower part of a pretrained network is a feature extractor that would be expensive to learn again and is probably right for your data. The upper part encodes the source task, so it is more likely to need changing. The final classification layer is always wrong for your task, because it predicts the source labels, so it is always replaced.
Four strategies, from least to most change
| Strategy | What trains | Cost | Risk |
|---|---|---|---|
| Linear probe (feature extraction) | new head only | lowest; features can be cached | underfits if the target needs new features |
| Partial fine-tune | head plus the last one or two blocks | moderate | overfits with very little data |
| Full fine-tune | every layer, small learning rate | highest memory and compute | destroys features if the learning rate is too high |
| Linear probe, then fine-tune | head first, then more layers | a little more than fine-tuning | the most robust default |
The last row deserves explanation. When you attach a randomly initialised head and fine-tune everything at once, the head's early gradients are large and essentially random, and they flow back into the pretrained body and distort features that were good. Kumar and colleagues (ICLR 2022) showed that full fine-tuning can underperform a linear probe on out-of-distribution data for this reason, and that training the head first and then fine-tuning (LP-FT) gets the in-distribution gains of fine-tuning while keeping more of the robustness. Howard and Ruder's ULMFiT (2018) reached a similar practice for text with gradual unfreezing: train the top first and unfreeze downwards.
Choosing a strategy from your data
| Target similar to source | Target different from source | |
|---|---|---|
| Small dataset (hundreds to a few thousand) | linear probe; maybe the last block | probe from a middle layer, or find a closer pretrained model |
| Large dataset (tens of thousands and up) | LP-FT, or full fine-tune | full fine-tune; consider domain-specific pretraining |
Similar means the inputs look alike: natural photographs to natural photographs, English web text to English support tickets. Different means a shift in what the input is: ImageNet features applied to grayscale X-rays, satellite multispectral bands or spectrograms. For a strong shift, the best move is often not a cleverer fine-tuning schedule but a better starting point: a model pretrained on data that resembles yours, if one with a usable licence exists.
Do not decide by reasoning alone. A linear probe takes minutes and gives you a floor; a short fine-tuning run tells you how much headroom there is. Everything else is tuning between those two numbers.
A worked example: six defect classes from 4,000 images
Suppose a factory has 4,000 labelled photographs of circuit boards in six classes: good, and five kinds of defect. Photographs of physical objects are close enough to ImageNet that the lower layers should transfer well, but defects are fine-grained, so the top block will probably need to adapt. That points to LP-FT with a partial unfreeze. The recipe below uses torchvision's weights API, which also gives you the exact preprocessing the weights were trained with.
import torch, torch.nn as nn
from torchvision.models import resnet50, ResNet50_Weights
weights = ResNet50_Weights.IMAGENET1K_V2
model = resnet50(weights=weights)
preprocess = weights.transforms() # the resize, crop and normalisation it was trained with
num_classes = 6
model.fc = nn.Linear(model.fc.in_features, num_classes) # 2048 -> 6, randomly initialised
def set_trainable(model, prefixes):
for name, p in model.named_parameters():
p.requires_grad = name.startswith(prefixes)
def freeze_bn_stats(model, trainable_prefixes):
# model.train() puts every BatchNorm back in training mode; undo it for frozen blocks.
for name, m in model.named_modules():
if isinstance(m, nn.BatchNorm2d) and not name.startswith(trainable_prefixes):
m.eval()
def run(model, loader, opt, sched, epochs, trainable):
loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)
for _ in range(epochs):
model.train(); freeze_bn_stats(model, trainable)
for x, y in loader:
opt.zero_grad(set_to_none=True)
loss = loss_fn(model(x.cuda()), y.cuda())
loss.backward(); opt.step(); sched.step()
model.cuda()
# Stage 1: linear probe. Only the new head learns.
stage1 = ("fc.",)
set_trainable(model, stage1)
opt = torch.optim.AdamW(model.fc.parameters(), lr=1e-3, weight_decay=1e-4)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=5 * len(train_loader))
run(model, train_loader, opt, sched, epochs=5, trainable=stage1)
# Stage 2: unfreeze the last block with a smaller learning rate than the head.
stage2 = ("fc.", "layer4.")
set_trainable(model, stage2)
opt = torch.optim.AdamW([
{"params": model.layer4.parameters(), "lr": 1e-4},
{"params": model.fc.parameters(), "lr": 3e-4},
], weight_decay=1e-4)
warm = 200 # linear warm-up steps, then cosine decay
sched = torch.optim.lr_scheduler.SequentialLR(opt, [
torch.optim.lr_scheduler.LinearLR(opt, start_factor=0.1, total_iters=warm),
torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=10 * len(train_loader) - warm),
], milestones=[warm])
run(model, train_loader, opt, sched, epochs=10, trainable=stage2)Read it as two experiments. Stage 1 trains 12,294 parameters (2,048 times 6 weights plus 6 biases) on top of a frozen 23.5-million-parameter body; it converges in a few epochs and gives a well-behaved head. Stage 2 lets layer4, about 15 million parameters in a ResNet-50, adapt at a learning rate a tenth of what a from-scratch run would use, while the head moves a little faster. Compare the validation score after each stage: if stage 2 adds nothing, the frozen features were enough and you have saved a deployment headache; if it adds a lot, try unfreezing layer3 too, with an even smaller rate.
Learning rates and regularisation
The single most common failure in fine-tuning is a learning rate that is too high for the pretrained layers. Pretrained weights sit in a good region of the loss surface; a large step jumps out of it, and the validation loss in the first epoch shows a spike that never fully recovers. Three habits prevent it.
- Discriminative learning rates. Give the head the largest rate and each lower unfrozen block a smaller one, as in the parameter groups above. A factor of two to ten between neighbouring groups is a common starting point.
- Warm-up. A short linear warm-up of a few hundred steps keeps the first updates gentle while the optimiser's moment estimates settle.
- Pull towards the start. Ordinary weight decay pulls weights towards zero, which is not where good pretrained weights are. L2-SP (Li and colleagues, 2018) instead penalises the distance from the pretrained values; with little data it is a cheap way to stop fine-tuning from wandering off.
Data augmentation matters more than usual, because the dataset is small; use augmentations that preserve the label for your task. A horizontal flip is harmless for most photos and wrong for text in images or for left-right medical findings.
BatchNorm: the detail that quietly breaks transfer
Convolutional networks such as ResNet contain BatchNorm layers, which keep running estimates of the mean and variance of their inputs, collected on the source data. Freezing a block's weights with requires_grad = False does not freeze those statistics: in training mode, every forward pass updates them from your batches. With small batches or a target distribution unlike the source, the statistics drift and the frozen layers no longer behave as they were trained.
Hence the freeze_bn_stats helper in the recipe. Calling model.train() switches every module to training mode, so frozen blocks have to be put back in evaluation mode after each call. For the blocks you do fine-tune, updating statistics is usually right if your batches have at least a few dozen examples; with very small batches, keep them frozen too, or use a model with GroupNorm or LayerNorm, which have no running statistics. Vision transformers use LayerNorm, so this particular trap does not apply to them.
Preprocessing is part of the model
A pretrained network expects inputs that look like what it saw in training: the same resolution, the same normalisation constants, the same channel order. Feed it images normalised with different statistics, or BGR instead of RGB, and accuracy drops for reasons that look like a modelling problem. In torchvision, weights.transforms() returns the inference preprocessing; build your training augmentation around the same final resize and normalisation. For text encoders the equivalent is the tokenizer, which must be the one shipped with the checkpoint, never a similar-looking one.
Resolution deserves a decision of its own. Fine-grained targets such as small defects often benefit from a higher input resolution than pretraining used. Convolutional networks accept it directly; vision transformers need their position embeddings interpolated, which the common libraries do for you but which is worth checking.
What the hardware does, and why feature caching wins
Training cost scales with what has to be stored for the backward pass. For each trainable parameter, mixed-precision training with Adam typically keeps the weight, a gradient and two optimiser moments, plus a full-precision master copy: on the order of 16 bytes. Full fine-tuning of ResNet-50's 25.6 million parameters therefore needs about 410 MB for parameters and optimiser state, before activations. Training only layer4 and the head needs that state for about 15 million parameters. Frozen parameters cost only their weights.
Activations follow the same logic with a twist: the backward pass needs stored activations only from the first trainable layer upwards. Freezing the lower blocks therefore saves activation memory and compute in the backward pass, not just optimiser state, which is why partial fine-tuning fits larger batches on the same GPU.
A linear probe goes further. If the body is frozen and you skip random augmentation, the features never change, so compute them once and train the head on the cached vectors. Each image costs one forward pass, ever; the head then trains on a CPU in seconds, and you can sweep regularisation strengths in a loop.
@torch.no_grad()
def embed(loader):
body = nn.Sequential(*list(model.children())[:-1]).eval().cuda() # everything but fc
feats, labels = [], []
for x, y in loader:
feats.append(body(x.cuda()).flatten(1).cpu()) # (batch, 2048)
labels.append(y)
return torch.cat(feats), torch.cat(labels)
X_train, y_train = embed(train_loader_no_augment)
X_val, y_val = embed(val_loader)
from sklearn.linear_model import LogisticRegression
probe = LogisticRegression(max_iter=2000, C=1.0).fit(X_train.numpy(), y_train.numpy())
print("linear probe accuracy:", probe.score(X_val.numpy(), y_val.numpy()))
Negative transfer and other failure modes
- Negative transfer. When source and target are too different, pretrained features can hurt: the fine-tuned model ends up worse than a small model trained from scratch. Always keep a from-scratch baseline for comparison if the domain is unusual.
- Catastrophic forgetting. Aggressive fine-tuning on a narrow dataset erases general features, and robustness to new conditions such as lighting or camera changes collapses even when the validation score looks fine. LP-FT and small learning rates are the defences.
- Leaked evaluation. If the pretraining data contains your test images, or near-duplicates of them, your scores are inflated. Public benchmarks are especially exposed.
- Overfitting the small set. Thousands of examples against millions of trainable parameters overfit quickly. Watch the gap between training and validation loss and stop early; gradient descent in depth covers reading those curves.
- Licence surprises. Pretrained weights carry licences, and some forbid commercial use. Check before the model ships, not after.
Evaluating honestly
Report three numbers side by side: the linear probe, your fine-tuned model, and, where feasible, a model trained from scratch. The probe tells you how good the borrowed features are; the gap to fine-tuning tells you what adaptation bought; the from-scratch baseline tells you whether transfer helped at all. Hold out data from the conditions you will actually deploy into, such as a different production line or a later month, because in-distribution validation hides the robustness differences between strategies. To understand how pretraining scale changes what fine-tuning can achieve, see transfer learning scaling; if the goal is a smaller deployable model rather than a new task, model distillation is the related tool.
What to do next
- Pick a pretrained model whose training data resembles yours and whose licence allows your use.
- Use the preprocessing shipped with the weights and confirm it with a quick sanity check on a few known images.
- Run a cached-feature linear probe and record the score as your floor.
- Fine-tune the last block with discriminative learning rates and a warm-up, keeping frozen BatchNorm layers in evaluation mode.
- Unfreeze further only while validation improves, lowering the rate for each deeper block.
- Evaluate on data from a shifted condition, not only the random split, and compare with a from-scratch baseline.
- Write down which layers trained, which rates you used and which pretrained checkpoint, so the result can be reproduced.