Training Vision Models

Domain adaptation for vision

A model trained on one camera, one lighting setup or one hospital falls apart on the next one, and domain adaptation is the set of tricks that recover the loss without new labels.

On this page 6
  1. What a domain gap looks like
  2. The cheapest fix that works
  3. When it is not enough
  4. Where you have already seen this
  5. Remember this
  6. 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.

Domain adaptation is fixing a model that works on your training photos and fails on the photos it actually meets.

Walk in from bright afternoon sunshine into a dim room. For a few seconds you see almost nothing. Then your eyes adjust and the furniture appears.

Nothing about the room changed. Nothing about your eyesight changed. Only the calibration changed, and it took two seconds.

An enormous share of "our model broke in production" is exactly that problem. The knowledge is intact. The calibration is wrong.

What a domain gap looks like

   TRAINED ON                    DEPLOYED ON
   bright shop lighting     ->   dim warehouse
   camera model A           ->   camera model B
   dry-season photos        ->   monsoon photos
   one hospital's scanner   ->   another hospital's scanner
   staged product shots     ->   phone photos from customers

              same objects, different look
                        |
                accuracy falls off a cliff

The word for this is domain shift. The pictures come from a different source than your training ones. The task has not changed.

You almost never get labels for the new domain. If somebody could label the warehouse photos, you would retrain and go home. The interesting problem is the one where you have plenty of new photos and no answers for them.

The cheapest fix that works

Here is a thing worth knowing about most vision models. Inside them are layers holding a running note of the data's average brightness and spread during training. Every later layer relies on that note.

Change the camera and the note is wrong. Everything downstream is subtly mis-scaled, and confident nonsense comes out.

The fix is to feed the model a pile of new photos, with no labels at all. It rewrites the note itself. That is it. No training, no gradients, no answers required.

In the developer section this takes a model from perfect, to worse than guessing, and back to perfect. Three lines of code.

When it is not enough

Recalibration handles changes in brightness, contrast and colour. It does nothing about changes in content.

If you trained on photos of packaged food and deploy on loose vegetables, no calibration saves you. That is not a domain gap, that is a different task. Being able to tell the two apart is the actual skill here.

Where you have already seen this

  • A phone camera that looks washed out for a moment when you step outside, then corrects itself.
  • A voice assistant that struggles in a new accent until it hears more of it.
  • A defect-detection line that needs a fresh check every time the factory changes its lighting.

Remember this

  • Domain shift means the same task, different-looking pictures.
  • Much of the damage comes from wrong calibration, not lost knowledge.
  • Recalibrating on unlabelled photos from the new domain is the first thing to try, and it is nearly free.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch numpy

Written and run against torch 2.13.0 (CPU) and numpy 2.2.6. No downloads. The whole script trains a small CNN from scratch on synthetic shapes and runs in a few seconds.

A camera change, and a three-line repair

adabn.py
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.data import TensorDataset, DataLoader

rng = np.random.default_rng(0); torch.manual_seed(0)

def draw(cls, bright, contrast):
    img = rng.normal(0.35, 0.06, (32, 32))
    r, c = rng.integers(9, 23, 2); s = int(rng.integers(4, 7))
    yy, xx = np.ogrid[:32, :32]; d = np.hypot(yy - r, xx - c)
    m = [(abs(yy-r) < s) & (abs(xx-c) < s), d < s, abs(d-s) < 1.2,
         (abs(yy-r) < 1.2) | (abs(xx-c) < 1.2)][cls]
    img[m] = 0.95
    return np.clip((img - 0.5) * contrast + 0.5 + bright, 0, 1).astype("float32")

def make(n, bright=0.0, contrast=1.0):
    lab = rng.integers(0, 4, n)
    x = torch.tensor(np.stack([draw(c, bright, contrast) for c in lab])).unsqueeze(1)
    return x, torch.tensor(lab)

# SOURCE = the camera you trained on.  TARGET = a dimmer, flatter camera.
Xs, ys = make(1200)
Xs_te, ys_te = make(400)
Xt, yt = make(400, bright=-0.18, contrast=0.45)
Xt_pool, _ = make(400, bright=-0.18, contrast=0.45)     # unlabelled target images

net = nn.Sequential(
    nn.Conv2d(1,16,3,padding=1), nn.BatchNorm2d(16), nn.ReLU(), nn.MaxPool2d(2),
    nn.Conv2d(16,32,3,padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2),
    nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(32,4))
opt = torch.optim.AdamW(net.parameters(), lr=3e-3)
dl = DataLoader(TensorDataset(Xs, ys), batch_size=64, shuffle=True)
for _ in range(25):
    net.train()
    for xb, yb in dl:
        opt.zero_grad(); F.cross_entropy(net(xb), yb).backward(); opt.step()

def acc(x, y):
    net.eval()
    with torch.no_grad():
        return (net(x).argmax(1) == y).float().mean().item()

print(f"same camera as training  : {acc(Xs_te, ys_te):.3f}")
print(f"different camera         : {acc(Xt, yt):.3f}")

# AdaBN: recompute BatchNorm statistics on UNLABELLED target images. No labels, no gradients.
for mod in net.modules():
    if isinstance(mod, nn.BatchNorm2d):
        mod.reset_running_stats(); mod.momentum = None    # None = plain running average
net.train()
with torch.no_grad():
    for i in range(0, 400, 64):
        net(Xt_pool[i:i+64])
print(f"different camera + AdaBN : {acc(Xt, yt):.3f}")
Output
same camera as training  : 1.000
different camera         : 0.275
different camera + AdaBN : 1.000

Deterministic across runs with the seeds set. Your numbers should match. If the middle line is not far below the first, check the brightness and contrast shift.

What that run shows

1.000 to 0.275 is a collapse, and the shift was mild. The target images are dimmer and flatter — brightness down 0.18, contrast scaled to 0.45. A human sees the same four shapes without effort. Four-class chance is 0.250, so the model finished barely above guessing.

Zero labels were used to fix it. Xt_pool was created, its labels discarded, and only the images passed through. This matters commercially: you always have unlabelled production images, and you rarely have labelled ones.

No weight moved. torch.no_grad() was on, no optimiser ran. Only the BatchNorm buffers changed. The classifier that scored 0.275 and the classifier that scored 1.000 have identical weights.

That is the point this lesson exists to make. The knowledge was never lost. It was being read through a wrong scale.

Line by line

mod.reset_running_stats() clears running_mean and running_var and sets num_batches_tracked to zero.

mod.momentum = None switches BatchNorm from an exponential moving average to a cumulative average over every batch it sees. With the default momentum of 0.1, the final statistics would be dominated by the last few batches. Setting it to None uses all 400 target images equally.

net.train() is required and looks wrong. Training mode is what makes BatchNorm update its buffers. See freezing and unfreezing layers for the same mechanism appearing as a bug. Here it is the feature.

acc() calls net.eval() internally, so measurement uses the newly stored statistics rather than the test batch's own.

When AdaBN will not save you

ShiftAdaBN helps?
Brightness, contrast, colour cast, exposureYes, often dramatically
Sensor noise, mild blur, compressionOften
New viewpoint or scalePartly
New object classes, new backgrounds, new taskNo

There is also a hard requirement: your architecture must contain BatchNorm. Vision transformers use LayerNorm. Its statistics are computed per sample at inference, so they need no adaptation and get no free repair.

Common mistakes

Adapting on a handful of target images. BatchNorm statistics estimated from 30 images are noisy, and you can end up worse than before. Fix: use several hundred, and check on any labels you do have.

Adapting on a target batch that is one class. If the production stream is sorted, a batch of 64 identical objects gives statistics for that object, not the domain. Fix: shuffle across a large buffer before adapting.

Adapting once and forgetting. Domains keep moving — seasons, lens dirt, a replaced bulb. Fix: recalibrate on a schedule and log the statistics so you can see drift. See monitoring and drift.

Calling it domain adaptation when the labels changed. If the new site also defines "defective" differently, no unsupervised method fixes that. Fix: label a small target set and measure, before assuming the problem is visual.

Try it yourself

Change contrast=0.45 to contrast=1.0 and keep only the brightness shift. Measure how much of the drop AdaBN recovers. Then shift the shapes' sizes instead of their brightness, and watch AdaBN stop helping. That is a content change wearing a domain-shift costume.

What to learn next

Researcher — Mathematics and papers.

The formal setting

Unsupervised domain adaptation assumes a labelled source $\mathcal{D}_S = {(x_i^s, y_i^s)}$ drawn from $p_S(x, y)$ and an unlabelled target $\mathcal{D}_T = {x_j^t}$ drawn from $p_T(x)$, with

$$ p_S(x) \neq p_T(x), \qquad p_S(y \mid x) = p_T(y \mid x) $$

  • $p_S, p_T$ — source and target distributions.
  • The second equality is the covariate shift assumption: the labelling rule is unchanged, only the inputs move.

Every method in this lesson rests on that second equality. It fails when the target site labels things differently. That is concept shift, and unsupervised adaptation is then unsound rather than weak.

Ben-David et al. (2010) give the bound that frames the field:

$$ \epsilon_T(h) \le \epsilon_S(h) + d_{\mathcal{H}\Delta\mathcal{H}}(\mathcal{D}_S, \mathcal{D}_T) + \lambda $$

  • $\epsilon_S, \epsilon_T$ — source and target risk of hypothesis $h$.
  • $d_{\mathcal{H}\Delta\mathcal{H}}$ — a divergence measuring how distinguishable the two domains are to the hypothesis class.
  • $\lambda$ — the error of the best joint hypothesis, which no algorithm controls.

Almost every deep method minimises the middle term. That $\lambda$ sits outside anyone's control is the theoretical reason adaptation sometimes cannot work.

The families

FamilyMechanismRepresentative
Statistic matchingAlign feature momentsSun and Saenko, 2016, Deep CORAL
Adversarial alignmentA domain discriminator through a gradient reversal layerGanin and Lempitsky, 2015, DANN
NormalisationRecompute BatchNorm statistics on targetLi et al., 2017, AdaBN
Pixel-levelTranslate source images into target styleHoffman et al., 2018, CyCADA
Test-timeAdapt online during inferenceWang et al., 2021, Tent; Sun et al., 2020, TTT

DANN inserts a gradient reversal layer. The feature extractor then maximises domain-classification loss while the domain classifier minimises it. The result is features from which the domain cannot be read. AdaBN is the outlier in this table. It is parameter-free, needing no joint training, no extra loss and no source data.

Test-time adaptation, and its failure modes

Tent (Wang et al., ICLR 2021) extends AdaBN by also updating the BatchNorm affine parameters $\gamma, \beta$ to minimise prediction entropy

$$ H(\hat{y}) = -\sum_{c} \hat{y}_c \log \hat{y}_c $$

  • $\hat{y}_c$ — predicted probability for class $c$ on a target batch.

Entropy is a proxy for error and needs no labels, so adaptation runs online, batch by batch, during deployment.

The proxy has a degenerate optimum, and this matters in production. Predicting one class for everything achieves zero entropy. Niu et al. (2023), Towards Stable Test-Time Adaptation in Dynamic Wild World (ICLR), document exactly this collapse. It appears under mixed shifts, small batches and imbalanced online label distributions. They stabilise it by discarding high-gradient noisy samples and seeking flat minima. If you deploy entropy-based adaptation, monitor prediction diversity as a first-class metric, not accuracy alone.

Papers

What to learn next