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.
- 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.
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 cliffThe 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
- Geometric augmentations — building tolerance to shift into training, before you need to adapt.
- BatchNorm in PyTorch — the buffers this whole lesson turns on.
- Monitoring and drift — noticing a domain shift before your users do.
Developer — Code and libraries.
Setup
pip install torch numpyWritten 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
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}")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
| Shift | AdaBN helps? |
|---|---|
| Brightness, contrast, colour cast, exposure | Yes, often dramatically |
| Sensor noise, mild blur, compression | Often |
| New viewpoint or scale | Partly |
| New object classes, new backgrounds, new task | No |
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
- Geometric augmentations — building tolerance to shift into training, before you need to adapt.
- BatchNorm in PyTorch — the buffers this whole lesson turns on.
- Monitoring and drift — noticing a domain shift before your users do.
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
| Family | Mechanism | Representative |
|---|---|---|
| Statistic matching | Align feature moments | Sun and Saenko, 2016, Deep CORAL |
| Adversarial alignment | A domain discriminator through a gradient reversal layer | Ganin and Lempitsky, 2015, DANN |
| Normalisation | Recompute BatchNorm statistics on target | Li et al., 2017, AdaBN |
| Pixel-level | Translate source images into target style | Hoffman et al., 2018, CyCADA |
| Test-time | Adapt online during inference | Wang 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
- Ben-David et al., A theory of learning from different domains, Machine Learning 2010
- Ganin and Lempitsky, Unsupervised Domain Adaptation by Backpropagation, 2015 — arxiv.org/abs/1409.7495
- Sun and Saenko, Deep CORAL, 2016 — arxiv.org/abs/1607.01719
- Li et al., Revisiting Batch Normalization For Practical Domain Adaptation (AdaBN), 2017 — arxiv.org/abs/1603.04779
- Wang et al., Tent: Fully Test-Time Adaptation by Entropy Minimization, 2021 — arxiv.org/abs/2006.10726
- Niu et al., Towards Stable Test-Time Adaptation in Dynamic Wild World, 2023 — arxiv.org/abs/2302.12400
What to learn next
- Geometric augmentations — building tolerance to shift into training, before you need to adapt.
- BatchNorm in PyTorch — the buffers this whole lesson turns on.
- Monitoring and drift — noticing a domain shift before your users do.