Transfer Learning and Pretraining

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.

The unfreezing pyramid, generic at the bottomStem and block 1 — edges: keep frozen longestBlock 2 — textures: unfreeze third, tiny rateBlock 3 — parts: unfreeze secondBlock 4 — objects: unfreezefirst, it is least genericNew head — random: train this alone to begin with
A random head sends a huge gradient backwards, so training it alone first is what stops it wrecking the pretrained blocks.

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,0982048 \times 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

Python
import torchimport torch.nn as nnfrom torchvision import modelsfrom torchvision.models import ResNet50_Weights# Load pretrained weights. The `weights=` enum is the current API;# the old `pretrained=True` flag is deprecated and prints a warning.weights = ResNet50_Weights.IMAGENET1K_V2model = models.resnet50(weights=weights)# 1. Freeze every parameter in the network.for param in model.parameters():    param.requires_grad = False# 2. Replace the head. Newly created layers have requires_grad=True#    by default, so this un-freezes exactly what we want.num_features = model.fc.in_features          # 2048 for ResNet-50model.fc = nn.Linear(num_features, 2)# 3. Give the optimiser ONLY the trainable parameters. Passing#    model.parameters() also works (parameters with no gradient are#    skipped), but this makes the intent explicit and easy to count.trainable = [p for p in model.parameters() if p.requires_grad]print(f"trainable: {sum(p.numel() for p in trainable):,}")   # 4,098optimiser = torch.optim.AdamW(trainable, lr=1e-3, weight_decay=1e-4)criterion = 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:

Python
model.train()          # sets everything to training mode# ...then force the frozen backbone back to eval so BatchNorm# uses its ImageNet running statistics and stops updating them.for module in model.modules():    if isinstance(module, nn.BatchNorm2d):        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

Python
import tensorflow as tffrom tensorflow.keras import layers, Modelbase = tf.keras.applications.ResNet50(    weights="imagenet",    include_top=False,        # drop the 1000-way classifier    input_shape=(224, 224, 3),)base.trainable = False        # freezes weights AND puts BN in inference modeinputs = tf.keras.Input(shape=(224, 224, 3))x = tf.keras.applications.resnet50.preprocess_input(inputs)x = base(x, training=False)   # belt-and-braces: keep BN in inference modex = layers.GlobalAveragePooling2D()(x)x = layers.Dropout(0.3)(x)outputs = layers.Dense(2, activation="softmax")(x)model = Model(inputs, outputs)model.compile(optimizer=tf.keras.optimizers.Adam(1e-3),              loss="sparse_categorical_crossentropy",              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.

Python
import numpy as npbackbone = nn.Sequential(*list(model.children())[:-1]).eval().cuda()feats, labels = [], []with torch.no_grad():    for images, y in loader:                 # ONE pass over the data        f = backbone(images.cuda()).flatten(1)   # (B, 2048)        feats.append(f.cpu()); labels.append(y)X = torch.cat(feats).numpy()                 # (420, 2048)y = torch.cat(labels).numpy()from sklearn.linear_model import LogisticRegressionclf = 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 extractionDisadvantages
Very fast — minutes, sometimes seconds with cached featuresAccuracy ceiling is lower when the target domain differs from the source
Almost impossible to overfit with a linear head on 2048 featuresCannot 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 GPUSensitive to input preprocessing mismatches, with no ability to compensate
Highly reproducible — very few moving partsWastes capacity when you do have enough data to fine-tune
Gives a trustworthy baseline that fine-tuning must beatFrozen 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\ln(2) \approx 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

Python
for param in model.parameters():    param.requires_grad = Trueoptimiser = 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.

Python
def freeze_all_but(model, groups_to_train):    """groups_to_train: list of nn.Module attributes to unfreeze."""    for p in model.parameters():        p.requires_grad = False    for g in groups_to_train:        for p in g.parameters():            p.requires_grad = Trueschedule = [    # (epochs, unfrozen groups,                    lr)    (3,  [model.fc],                               1e-3),    (4,  [model.fc, model.layer4],                 1e-4),    (4,  [model.fc, model.layer4, model.layer3],   5e-5),]for n_epochs, groups, lr in schedule:    freeze_all_but(model, groups)    opt = torch.optim.AdamW(        [p for p in model.parameters() if p.requires_grad],        lr=lr, weight_decay=1e-4)    for _ in range(n_epochs):        train_one_epoch(model, loader, opt, criterion)    validate(model, val_loader)   # check before widening further

The 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 classWhat to unfreezeApprox. trainable params (ResNet-50)Typical LR
Under 100Head only~4K–20K1e-3
100–500Head + layer4~15M1e-4
500–2,000Head + layer4 + layer3~22M5e-5 to 1e-4
2,000–10,000Everything except conv1 and layer1~25M3e-5 to 5e-5
Over 10,000Everything25.6M1e-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.

Python
base_lr = 1e-3factor = 2.6groups = [model.fc, model.layer4, model.layer3, model.layer2, model.layer1]param_groups = []for depth, group in enumerate(groups):    param_groups.append({        "params": group.parameters(),        "lr": base_lr / (factor ** depth),    })optimiser = torch.optim.AdamW(param_groups, weight_decay=1e-4)

The resulting rates, computed exactly:

GroupFormulaLearning rateRelative to head
fc (new head)10−3/2.6010^{-3} / 2.6^{0}1.00e-31×
layer410−3/2.6110^{-3} / 2.6^{1}3.85e-40.38×
layer310−3/2.6210^{-3} / 2.6^{2}1.48e-40.15×
layer210−3/2.6310^{-3} / 2.6^{3}5.69e-50.057×
layer110−3/2.6410^{-3} / 2.6^{4}2.19e-50.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

DimensionFeature extractionFine-tuning
Trainable parametersThousandsMillions to tens of millions
Data requiredTens per class is workableHundreds to thousands per class
Training time (420 images)Seconds to a few minutesTens of minutes to hours
GPU memoryLow — no backbone gradients or optimiser stateHigh — gradients plus optimiser state for every weight
Overfitting riskVery lowHigh without regularisation
Accuracy ceilingLimited by how well source features describe your dataSubstantially higher when data supports it
Sensitivity to learning rateForgivingUnforgiving — the single biggest cause of failure
Handles large domain shiftPoorlyWell, given enough data
ReproducibilityExcellentNoticeable 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 imagesFrom scratchFeature extractionFine-tune layer4+headFull fine-tune
20056.5%96.8%96.1%91.2%
2,00072.3%97.4%98.6%98.1%
20,00091.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

Text
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

SettingTypical LR (SGD)Typical LR (AdamW)
Training ResNet from scratch0.11e-3
New head only, backbone frozen1e-21e-3
Fine-tuning the top block1e-31e-4
Full fine-tuning1e-41e-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

Python
from torchvision import transformstrain_tf = transforms.Compose([    transforms.RandomResizedCrop(224, scale=(0.7, 1.0)),    transforms.RandomHorizontalFlip(),    transforms.RandomRotation(15),    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),    transforms.ToTensor(),    transforms.Normalize([0.485, 0.456, 0.406],                         [0.229, 0.224, 0.225]),    transforms.RandomErasing(p=0.25, scale=(0.02, 0.15)),])

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

Python
model.fc = nn.Sequential(    nn.Dropout(0.5),    nn.Linear(2048, 512),    nn.BatchNorm1d(512),    nn.ReLU(inplace=True),    nn.Dropout(0.3),    nn.Linear(512, num_classes),)

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.

Python
best_acc, patience, bad_epochs = 0.0, 5, 0for epoch in range(30):    train_one_epoch(model, train_loader, optimiser, criterion)    acc = validate(model, val_loader)    if acc > best_acc:        best_acc, bad_epochs = acc, 0        torch.save(model.state_dict(), "best.pt")    else:        bad_epochs += 1        if bad_epochs >= patience:            breakmodel.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.

Python
import torch, torch.nn as nnfrom torchvision import modelsfrom torchvision.models import ResNet50_Weightsmodel = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)model.fc = nn.Sequential(nn.Dropout(0.4), nn.Linear(2048, 2))model = model.cuda()criterion = nn.CrossEntropyLoss(label_smoothing=0.05)def set_trainable(groups):    for p in model.parameters():        p.requires_grad = False    for g in groups:        for p in g.parameters():            p.requires_grad = Truedef make_opt(lr):    return torch.optim.AdamW(        [p for p in model.parameters() if p.requires_grad],        lr=lr, weight_decay=1e-4)# Phase 1 - head only. Establishes the baseline AND warms up the head# so that phase 2 does not blow up the pretrained features.set_trainable([model.fc])run(epochs=5, opt=make_opt(1e-3))       # expect ~84%# Phase 2 - add the last residual block at a 10x lower rate.set_trainable([model.fc, model.layer4])run(epochs=8, opt=make_opt(1e-4))       # expect ~89%# Phase 3 - discriminative rates across the whole network.for p in model.parameters():    p.requires_grad = Trueopt = torch.optim.AdamW([    {"params": model.fc.parameters(),     "lr": 1e-4},    {"params": model.layer4.parameters(), "lr": 3.8e-5},    {"params": model.layer3.parameters(), "lr": 1.5e-5},    {"params": model.layer2.parameters(), "lr": 5.7e-6},    {"params": model.layer1.parameters(), "lr": 2.2e-6},], weight_decay=1e-4)run(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:

  1. 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.
  2. Frozen backbone, small MLP head with dropout. Occasionally worth a point or two when class boundaries are not linearly separable in feature space.
  3. Unfreeze the last block, learning rate 10× lower. The single highest-value step for most projects.
  4. 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.