Training Vision Models

Linear probing vs fine-tuning

Linear probing freezes the backbone and trains one small layer on top, while fine-tuning moves everything, and the choice decides your accuracy, your cost and how well the model survives new data.

On this page 5
  1. Why the choice matters
  2. What you actually get
  3. Where you have already seen this
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Linear probing trains only the last small layer. Fine-tuning trains the whole model.

Imagine a friend on a phone call describing each photo in your album, one fixed sentence each. "A dark round thing on a metal sheet." Your job is to sort the photos into piles from those sentences alone.

Linear probing is learning to sort the sentences. Your friend never changes how she describes things.

Fine-tuning is also retraining your friend, so her sentences start mentioning the details your sorting job cares about.

Why the choice matters

Both routes start from the same borrowed model. They differ in what is allowed to move.

LINEAR PROBING
  photo -> [ backbone: LOCKED ] -> 512 numbers -> [ small new layer ] -> answer
                                                    only this layer trains

FINE-TUNING
  photo -> [ backbone: OPEN   ] -> 512 numbers -> [ small new layer ] -> answer
             this trains as well

Because the backbone is locked in probing, the 512 numbers for a given photo never change. So you compute them once, store them, and then train on stored numbers. Training becomes as cheap as sorting a spreadsheet.

In fine-tuning nothing can be stored, because every weight moves after every batch. Each epoch means running every photo through the whole model again, forwards and backwards.

What you actually get

Fine-tuning usually wins on accuracy, given enough data and enough time. Probing usually wins on everything else: speed, memory, stability, and the ability to run on a laptop.

There is one more difference that costs teams real money. Fine-tuning moves the backbone towards your particular photos. That is helpful when the new photos look like the old ones, and harmful when they do not. Probing cannot overfit the backbone, because it never touches it.

The recipe most people land on: probe first, read the number, fine-tune only if it is not good enough.

Where you have already seen this

Photo apps that let you tag one person and then find them everywhere are usually probing. There is no time to retrain a vision model on your phone while you wait.

A hospital tool tuned over months on a hundred thousand scans is usually fine-tuned. The data and the budget are both there.

Remember this

  • Probing locks the backbone and trains one small layer, so features can be computed once.
  • Fine-tuning moves every weight, usually scoring higher and costing far more.
  • Probe first. It is often close enough, and it tells you whether the features suit your task at all.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch torchvision scikit-learn

Written and run against torch 2.13.0 (CPU), torchvision 0.28.0, scikit-learn 1.7.2.

This example uses a real dataset: fetch_olivetti_faces from scikit-learn. It holds 400 photographs of 40 people, 64 x 64 greyscale, about 4.5 MB. Plus the 44.7 MB ResNet-18 checkpoint. Total run time on a laptop CPU was under twenty seconds.

Both methods, same backbone, same split

probe_vs_finetune.py
import time, torch, torch.nn as nn, torch.nn.functional as F, numpy as np
from sklearn.datasets import fetch_olivetti_faces
from sklearn.linear_model import LogisticRegression
from torch.utils.data import TensorDataset, DataLoader
from torchvision.models import resnet18, ResNet18_Weights

torch.manual_seed(0); np.random.seed(0)

faces = fetch_olivetti_faces()                    # 400 real photos of 40 people, ~4.5 MB
x = torch.tensor(faces.images).unsqueeze(1).repeat(1, 3, 1, 1)   # grey -> 3 channels
x = F.interpolate(x, size=112, mode="bilinear", align_corners=False)
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, y = (x - mean) / std, torch.tensor(faces.target)

tr = torch.cat([torch.arange(i*10, i*10+6) for i in range(40)])     # 6 per person to learn from
te = torch.cat([torch.arange(i*10+6, i*10+10) for i in range(40)])  # 4 per person to test on

def test_acc(model):
    model.eval()
    with torch.no_grad():
        pred = torch.cat([model(x[te][i:i+40]) for i in range(0, len(te), 40)]).argmax(1)
    return (pred == y[te]).float().mean().item()

# ---------- linear probe: freeze everything, fit one linear layer ----------
t0 = time.time()
backbone = resnet18(weights=ResNet18_Weights.DEFAULT)
backbone.fc = nn.Identity()                       # stop at the 512-number description
backbone.eval()
with torch.no_grad():
    feats = torch.cat([backbone(x[i:i+40]) for i in range(0, 400, 40)])
probe = LogisticRegression(max_iter=2000).fit(feats[tr], y[tr])
print(f"linear probe   accuracy {probe.score(feats[te], y[te]):.3f}   {time.time()-t0:.0f}s total")

# ---------- fine-tune: same backbone, every weight now moves ----------
t0 = time.time()
model = resnet18(weights=ResNet18_Weights.DEFAULT)
model.fc = nn.Linear(512, 40)
opt = torch.optim.AdamW(model.parameters(), lr=1e-4)
loader = DataLoader(TensorDataset(x[tr], y[tr]), batch_size=16, shuffle=True)
for ep in range(4):
    model.train()
    for xb, yb in loader:
        opt.zero_grad()
        F.cross_entropy(model(xb), yb).backward()
        opt.step()
    print(f"fine-tune ep{ep+1} accuracy {test_acc(model):.3f}   {time.time()-t0:.0f}s so far")
Output
linear probe   accuracy 0.887   2s total
fine-tune ep1 accuracy 0.237   3s so far
fine-tune ep2 accuracy 0.613   5s so far
fine-tune ep3 accuracy 0.900   9s so far
fine-tune ep4 accuracy 0.969   13s so far

The accuracies repeated exactly across runs on one machine with the seeds set. The timings will not match yours — they depend on your CPU, thread count and background load. Repeated runs on the same laptop finished the fine-tune anywhere between 8 and 14 seconds. Read the ratio, not the numbers: the probe cost roughly a sixth of the fine-tune here.

What this run is telling you

The probe hit 0.887 in two seconds, and never trained the network. Forty-way face identification, six examples per person. The backbone was trained on internet objects, never on faces. Random guessing here is 0.025.

Fine-tuning was worse than the probe for the first two epochs. 0.237, then 0.613. The head is random at the start. Its early gradients push the backbone around before carrying any useful signal. That is not a hyperparameter accident. It is the mechanism described by Kumar et al. in the researcher section.

Fine-tuning passed the probe at epoch three and finished at 0.969. Eight extra points of accuracy for roughly six times the compute, on 240 training images. Whether that trade is worth it is a product question, not a machine learning one.

The probe's real advantage is hidden in the code. Features were computed once, for all 400 images, then reused. Add ten more people and you extract features for the new photos only. Refit a logistic regression in a second, and ship. The fine-tuned model has to be retrained from the checkpoint.

Line by line

backbone.fc = nn.Identity() replaces the classifier with a layer that returns its input unchanged. The model now outputs the 512-number description rather than class scores. This is cleaner than a forward hook and easier to read than slicing the model.

backbone.eval() before extraction is required, not optional. In train mode the BatchNorm layers rewrite their statistics as batches pass through. The "frozen" features would then depend on extraction order. See freezing and unfreezing layers.

LogisticRegression(max_iter=2000) is a linear classifier fitted by scikit-learn. It is the same maths as an nn.Linear head trained with cross-entropy. Its solver converges reliably on 240 rows, with no learning rate to pick. See logistic regression.

x[i:i+40] batching during extraction keeps peak memory low. 400 images at 112 x 112 through a ResNet-18 in one go is fine here. At 224 x 224 on a real dataset it is not.

Common mistakes

Fitting the probe on features extracted in train mode. The features drift, the probe fits the drift, and the score does not reproduce. Fix: backbone.eval() and torch.no_grad().

Judging fine-tuning by epoch one. As the output above shows, epoch one can be far below the probe. Fix: run at least until the validation curve flattens before declaring a winner.

Comparing a tuned probe against an untuned fine-tune. Probing has one important knob, the regularisation strength C. Fine-tuning has learning rate, schedule, epochs, augmentation and freezing policy. An unfair comparison is the norm, in both directions. Fix: state the budget you gave each.

Extracting features once and then augmenting. Cached features cannot be augmented — the augmentation happens on pixels, upstream of the cache. Fix: extract several augmented copies per image, or accept that probing gives up augmentation.

Try it yourself

Cut the training set from 6 photos per person to 2 by changing i*10+6 to i*10+2 in the tr line. Predict which method suffers more before you run it. The answer is in the next lesson but one, on few-shot classification.

What to learn next

Researcher — Mathematics and papers.

Cost, stated exactly

Let $n$ be the number of training images and $d$ the feature width. Let $C$ be the class count, $E$ the epochs, and $F$ the cost of one forward pass.

Feature costTraining costParameters updated
Linear probe$n F$ once$O(E \, n \, d \, C)$ on cached vectors$C(d+1)$
Fine-tuningnone cached$O(E \, n \cdot 3F)$all of $\phi$ plus $C(d+1)$

The factor 3 is the usual forward-plus-backward-plus-update accounting. The asymmetry is not marginal. For ResNet-18 at 224 x 224, $F \approx 1.8$ GFLOPs. A probe step on a 512-vector with 40 classes is about 20 kFLOPs. Probing is roughly five orders of magnitude cheaper per example after extraction.

Why fine-tuning can lose out of distribution

Kumar et al. (2022) (ICLR) analyse the overparameterised linear setting and prove the mechanism. Write the model as $w^{\top}\phi(x)$ with pretrained features $\phi$. Fine-tuning from a randomly initialised $w$ takes large early steps. They point wherever the random head happens to point. Those steps rotate $\phi$ within the subspace spanned by the training distribution. Directions that only matter out of distribution are left untouched. The feature geometry that generalised is distorted, and the distortion is invisible on in-distribution validation.

Their empirical result across 10 distribution-shift datasets is stark. Fine-tuning averages 2% better in-distribution and 7% worse out-of-distribution than probing.

The fix they propose is LP-FT — probe to convergence, then fine-tune from that head:

stage 1: freeze phi, fit w            -> w is now near-optimal for the frozen features
stage 2: unfreeze phi, small LR        -> early gradients are small, phi barely rotates

They report LP-FT beating full fine-tuning by roughly 1% in-distribution and 10% out-of-distribution. This is close to free, and it is under-used.

The weight-space alternative

Wortsman et al. (2022), Robust fine-tuning of zero-shot models (CVPR), take a different route. It applies to models with a usable zero-shot classifier, such as CLIP. Fine-tune normally, then interpolate in weight space:

$$ \theta_{\text{WiSE}} = \alpha \, \theta_{\text{fine-tuned}} + (1-\alpha)\, \theta_{\text{zero-shot}} $$

  • $\alpha \in [0,1]$ — the mixing coefficient, chosen on validation data.
  • $\theta_{\text{zero-shot}}$ — the weights before any fine-tuning.

On ImageNet and five derived shifts they report 4 to 6 points better accuracy under shift than prior work. ImageNet accuracy rose 1.6 points, at no extra training or inference cost. Averaging two sets of weights should not work as well as it does. That it does is evidence the two solutions lie in a connected low-loss region.

When probing wins outright

Two regimes, both common:

  • Very small $n$. With tens of examples per class, fine-tuning has more parameters than constraints and memorises. See few-shot image classification.
  • Frozen-backbone deployment. One extracted feature set can serve many downstream heads: search, dedup, clustering, several classifiers. That is only possible if the backbone never changes. Fine-tuning per task forks the representation and multiplies your serving cost.

Tian et al. (2020), Rethinking Few-Shot Image Classification: a Good Embedding Is All You Need? (ECCV), make the sharper version of this point. A plain linear classifier on a well-trained embedding outperformed the specialised meta-learning methods of the day. Representation quality dominated algorithm choice.

Papers

  • Kumar et al., Fine-Tuning can Distort Pretrained Features and Underperform Out-of-Distribution, 2022 — arxiv.org/abs/2202.10054
  • Wortsman et al., Robust fine-tuning of zero-shot models, 2022 — arxiv.org/abs/2109.01903
  • Tian et al., Rethinking Few-Shot Image Classification, 2020 — arxiv.org/abs/2003.11539
  • Chen et al., A Simple Framework for Contrastive Learning of Visual Representations (SimCLR, which made linear probing the standard evaluation), 2020 — arxiv.org/abs/2002.05709

What to learn next