Fine-tuning a vision model
Fine-tuning takes a model somebody else trained on millions of photos and continues its training gently on your few hundred, so you inherit their vision instead of paying for it.
- 11 min read
- 3 reading levels
- Updated
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Fine-tuning means taking a model somebody else already trained, and continuing its training on your own photos.
Think about hiring an experienced cook for your kitchen. She has spent fifteen years judging heat, holding a knife, knowing when oil is ready. You do not teach her any of that. You teach her your four recipes, and she is useful by the weekend.
Training a vision model from nothing is hiring somebody who has never entered a kitchen. Fine-tuning is hiring the experienced cook.
That is the whole idea. Everything else on this page is detail.
Why it exists
The models you read about were trained on enormous piles of photos. One common pile, called ImageNet, holds over a million labelled pictures. Training on it takes days of expensive machines.
You have eight hundred photos of four kinds of defect on a factory line. Nobody is going to give you a million.
Here is the part that makes fine-tuning work. Most of what a vision model learns is not about the specific objects it was shown. Early on it learns edges. Then textures. Then corners, curves, repeated patterns, the look of shiny metal against cloth.
Those skills are the same whether the photo shows a labrador or a cracked weld. Somebody already paid for them. You get to keep them.
How it works
a million internet photos
|
v
[ weeks of training, somebody else's bill ]
|
v
a model that understands edges, textures, shapes, objects
|
v
[ you continue training on YOUR 800 photos, gently ]
|
v
a model that knows YOUR four classesTwo things change when the model becomes yours.
The last layer is replaced. The original model ends in a layer that names one thousand internet categories. You throw that away and put in a fresh one that names your four. That layer knows nothing yet, which is fine, because it is small.
Every other layer moves a little. You keep training, but with tiny steps. Big steps would wipe out the fifteen years of kitchen experience you were trying to keep.
Where you have already seen this
- A plant-identification app that names a leaf from your camera.
- A hospital tool that flags which scans a doctor should open first.
- A factory camera that spots a cracked casting on a moving belt.
- A wildlife camera trap that sorts night photos into deer, boar and empty.
Almost none of these were trained from nothing. Almost all of them are fine-tuned.
What is honestly hard here
Fine-tuning looks like three lines of code and behaves like a temperamental machine. The most common failure is not a crash. It is a model that scores beautifully on your own photos and falls apart on next month's photos.
Eight hundred photos from one factory, one camera, one week carry patterns you never meant to teach. Read that twice. It is the single biggest reason fine-tuned models disappoint after they ship.
Remember this
- Fine-tuning continues somebody else's training instead of starting over.
- You replace the final layer and move every other layer by tiny amounts.
- The reused skills are edges, textures and shapes, which are shared across almost every visual task.
What to learn next
- Freezing and unfreezing layers — deciding which parts of the backbone are allowed to move.
- Transfer learning in PyTorch — the same idea from the PyTorch mechanics side.
- LoRA — updating a huge model by training a very small number of extra weights.
Developer — Code and libraries.
Setup
pip install torch torchvisionEverything below was written and run against torch 2.13.0 (CPU build) and torchvision 0.28.0 on a laptop CPU. The APIs used here have been stable since torchvision 0.13, when the weights= enum replaced pretrained=True.
Loading ResNet18_Weights.DEFAULT downloads a 44.7 MB checkpoint on first use and caches it. That is the only download this page needs.
A complete fine-tune, start to finish
import torch, torch.nn as nn
from torch.utils.data import TensorDataset, DataLoader
from torchvision.models import resnet18, ResNet18_Weights
torch.manual_seed(0)
# ---- a tiny fake dataset: 64 images, 64x64 RGB, two classes --------------
# class 0 = vertical stripes, class 1 = horizontal stripes, plus noise
def make(n, cls):
x = torch.rand(n, 3, 64, 64) * 0.3
if cls == 0:
x[:, :, :, ::4] += 0.7 # bright vertical lines
else:
x[:, :, ::4, :] += 0.7 # bright horizontal lines
return x.clamp(0, 1)
X = torch.cat([make(32, 0), make(32, 1)])
y = torch.cat([torch.zeros(32), torch.ones(32)]).long()
perm = torch.randperm(64)
X, y = X[perm], y[perm]
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
X = (X - mean) / std # the stats the weights were trained with
train = DataLoader(TensorDataset(X[:48], y[:48]), batch_size=8, shuffle=True)
Xte, yte = X[48:], y[48:]
# ---- the fine-tuning part ------------------------------------------------
model = resnet18(weights=ResNet18_Weights.DEFAULT)
model.fc = nn.Linear(model.fc.in_features, 2) # 1000 ImageNet classes -> 2 of ours
def accuracy():
model.eval()
with torch.no_grad():
return (model(Xte).argmax(1) == yte).float().mean().item()
print(f"accuracy before any training: {accuracy():.2f}")
opt = torch.optim.AdamW(model.parameters(), lr=1e-4) # small LR: do not wreck the features
loss_fn = nn.CrossEntropyLoss()
for epoch in range(3):
model.train()
total = 0.0
for xb, yb in train:
opt.zero_grad()
loss = loss_fn(model(xb), yb)
loss.backward()
opt.step()
total += loss.item() * len(xb)
print(f"epoch {epoch+1} train loss {total/48:.3f} held-out accuracy {accuracy():.2f}")accuracy before any training: 0.56 epoch 1 train loss 0.374 held-out accuracy 1.00 epoch 2 train loss 0.045 held-out accuracy 1.00 epoch 3 train loss 0.070 held-out accuracy 1.00
Those accuracies were identical across repeated runs on one machine with the seed set. Loss digits can shift on a different CPU, GPU or torch build. Floating-point reduction order is not guaranteed to match. Reproduce the shape of the run, not the digits.
Reading that output
0.56 before training is the sound of a random head. The fresh nn.Linear(512, 2) has never seen a gradient. On sixteen held-out images it got nine right, which is a coin toss. The backbone underneath is excellent and the head on top is noise.
One epoch was enough. Six batches of eight images moved a head from noise to a perfect score on this held-out set. That is the point of the whole technique: the features were already there, waiting to be pointed at something.
The task here is deliberately easy. Sixteen held-out images is a tiny test set, and stripes are a caricature of a real problem. Treat this script as a wiring diagram, not as evidence about accuracy.
Line by line, for the parts that bite
model.fc = nn.Linear(model.fc.in_features, 2) reads in_features from the layer being replaced rather than typing 512. Switch to resnet50 and the number becomes 2048 with no edit. See replacing the classifier head for architectures that do not call it fc.
(X - mean) / std uses ImageNet's channel statistics. The pretrained weights were trained on inputs normalised this way, and feeding raw 0-to-1 pixels shifts every activation. Accuracy often survives, then quietly disappoints.
lr=1e-4 is roughly ten to a hundred times smaller than a from-scratch learning rate. The backbone holds knowledge you paid nothing for and want to keep.
model.train() and model.eval() are not decoration. ResNet contains BatchNorm layers that behave differently in each mode, and other architectures add dropout. See train and eval mode.
Common mistakes
Using a from-scratch learning rate. At lr=1e-2 the first gradients from a random head flood backwards. They scramble the pretrained weights before the head has learned anything. Symptom: training loss falls, held-out accuracy stays near chance. Fix: start at 1e-4, or freeze the backbone for one epoch first.
Skipping normalisation, or using the wrong statistics. Fix: ResNet18_Weights.DEFAULT.transforms() returns the exact preprocessing the checkpoint expects, including resize, crop and normalisation. Use it rather than copying numbers by hand.
Leaving the model in train() while measuring. BatchNorm then updates its running statistics from your test batch, so the score depends on batch size and ordering. Fix: model.eval() plus torch.no_grad() around every evaluation.
Testing on photos that share a source with training. Two crops of one photo, split across train and test, give a score that means nothing. Two frames from one video do the same. Fix: split by source — by patient, by camera, by day — not by file. See train-test split.
Try it yourself
Change lr=1e-4 to lr=1e-1 and run it again. Watch the loss go somewhere strange and the accuracy collapse. You have destroyed a million-photo education in six batches. That failure is worth seeing once, on purpose, on a script you do not care about.
What to learn next
- Freezing and unfreezing layers — deciding which parts of the backbone are allowed to move.
- Transfer learning in PyTorch — the same idea from the PyTorch mechanics side.
- LoRA — updating a huge model by training a very small number of extra weights.
Researcher — Mathematics and papers.
What fine-tuning optimises
Let $\theta_0$ be the pretrained parameters and $\theta = (\phi, w)$ the fine-tuned pair of backbone $\phi$ and head $w$. Standard fine-tuning minimises
$$ \mathcal{L}(\theta) = \frac{1}{n}\sum_{i=1}^{n} \ell\big(f_\theta(x_i),\, y_i\big) + \lambda \lVert \theta \rVert_2^2 $$
- $n$ — the number of labelled target images.
- $\ell$ — the per-example loss, usually cross-entropy.
- $f_\theta$ — the network mapping an image to class logits.
- $x_i, y_i$ — the $i$-th target image and its label.
- $\lambda$ — weight decay strength.
- $\lVert \cdot \rVert_2$ — the Euclidean norm.
Note what the objective does not contain: any term tying $\theta$ to $\theta_0$. Nothing in the loss asks the model to remember what it knew. The initialisation is the only inheritance, and gradient descent is free to walk away from it. Small learning rates, early stopping and weight decay act as implicit anchors. Explicit anchors exist too. $L^2\text{-}SP$ (Xuhong et al., 2018) replaces $\lVert\theta\rVert_2^2$ with $\lVert\theta - \theta_0\rVert_2^2$. That penalises distance from the pretrained point rather than from the origin.
Why features transfer, and how far
Yosinski et al. (2014), How transferable are features in deep neural networks? (NeurIPS), ran the canonical experiment. Split ImageNet in half, train on one half, transfer the first $k$ layers. Two findings still shape practice. Transferability decays with depth, because higher layers specialise to the source labels. And transferred-then-frozen layers can hurt through co-adaptation, when layers that learned to work together get separated. That effect is distinct from specialisation.
Kornblith, Shlens and Le (2019), Do Better ImageNet Models Transfer Better? (CVPR), evaluated 16 architectures across 12 datasets. ImageNet top-1 correlates strongly with transfer accuracy under fine-tuning. The relationship weakens on fine-grained datasets whose classes are absent from ImageNet. Several regularisers that help ImageNet also hurt transfer.
The learning-rate question
The head starts from an initialiser; the backbone starts from an optimum. One learning rate serves both badly. Discriminative learning rates give layer $l$ a rate
$$ \eta_l = \eta_{L} \cdot \gamma^{\,L-l} $$
- $\eta_L$ — the rate used for the final layer.
- $\gamma \in (0, 1]$ — the decay factor per layer, commonly $0.65$ to $0.9$.
- $L$ — the number of layers, counted from the input.
That is the ULMFiT recipe (Howard and Ruder, 2018, ACL), and it carries into vision unchanged. In PyTorch it is expressed with optimiser parameter groups. Ro and Choi (2021), AutoLR, search $\gamma$ automatically. They report that rates increasing towards the head are the useful family.
Catastrophic forgetting and the parameter-efficient alternatives
Full fine-tuning writes to every weight, so the source capability is lost. Each downstream task then needs its own copy of the network. Several families avoid that:
| Family | Trainable fraction | Representative work |
|---|---|---|
| Adapters | ~1-5% | Houlsby et al., 2019; Rebuffi et al., 2017 (residual adapters, vision) |
| Low-rank updates | ~0.1-1% | Hu et al., 2021, LoRA |
| Prompt / prefix | <1% | Jia et al., 2022, Visual Prompt Tuning (ECCV) |
| Bias-only | ~0.1% | Zaken et al., 2022, BitFit |
LoRA constrains the update to $\Delta W = BA$, with $B \in \mathbb{R}^{d \times r}$ and $A \in \mathbb{R}^{r \times k}$. Here $W \in \mathbb{R}^{d \times k}$ is the frozen matrix, and $d$ and $k$ are its output and input widths. The rank budget $r$ satisfies $r \ll \min(d,k)$. Jia et al. (2022) report visual prompt tuning beating full fine-tuning on 20 of 24 downstream tasks. The backbone was a pretrained ViT. That inverts the assumption that more trainable parameters win.
Papers
- Yosinski et al., How transferable are features in deep neural networks?, 2014 — arxiv.org/abs/1411.1792
- Kornblith et al., Do Better ImageNet Models Transfer Better?, 2019 — arxiv.org/abs/1805.08974
- Howard and Ruder, ULMFiT, 2018 — arxiv.org/abs/1801.06146
- Xuhong et al., Explicit Inductive Bias for Transfer Learning, 2018 — arxiv.org/abs/1802.01483
- Hu et al., LoRA, 2021 — arxiv.org/abs/2106.09685
- Jia et al., Visual Prompt Tuning, 2022 — arxiv.org/abs/2203.12119
What to learn next
- Freezing and unfreezing layers — deciding which parts of the backbone are allowed to move.
- Transfer learning in PyTorch — the same idea from the PyTorch mechanics side.
- LoRA — updating a huge model by training a very small number of extra weights.