Transfer Learning and Pretraining

Domain Adaptation — Closing the Train/Deploy Gap


A manufacturer deploys a surface-defect detector on production Line A. It was fine-tuned on 12,000 labelled images from that line's cameras and it reaches an F1 score of 94.2% on a held-out test set from the same line. Everyone is pleased. Six weeks later the same model is installed on Line B — identical product, identical defect types, identical camera model.

F1 on Line B: 68.5%.

Nothing is broken. The model has not degraded, the code is the same, the weights are byte-for-byte identical. The only difference is that Line B's overhead LEDs are a slightly cooler colour temperature and mounted 40 cm higher, so specular highlights fall in different places and the average image is about 18% brighter. To a human inspector the two lines look the same. To a network whose features were tuned on Line A's exact illumination, they are different worlds.

There are no labels for Line B. Getting 12,000 more images annotated by a metallurgist would take four months and cost more than the project. The question is what you can do with unlabelled target data, and the answer is a family of techniques called domain adaptation.

Line B breaks the model: fixes in order of costMeasure theshift beforefixing itRecomputeBatchNormstats on Line BSelf-train onconfidentpseudo-labelsAlign features:MMD or adversarialLabel a fewhundredtarget imagesFrozen BatchNorm normalises Line B images by Line A's mean and variance, which is often the whole failure.
The first fix needs no labels and no gradient steps — only a forward pass over unlabelled images from the new line.

What domain shift actually is

Standard supervised learning rests on one assumption that is almost never stated out loud: training and deployment data are drawn from the same distribution. When that fails, you have domain shift, and every guarantee your validation set gave you evaporates.

Write the source distribution as Ps(X,Y)P_s(X, Y) and the target as Pt(X,Y)P_t(X, Y). Domain shift means Ps(X,Y)≠Pt(X,Y)P_s(X, Y) \neq P_t(X, Y). But the joint distribution factors two ways, P(X,Y)=P(Y∣X)P(X)=P(X∣Y)P(Y)P(X,Y) = P(Y|X)P(X) = P(X|Y)P(Y), and which factor moved determines which fix works. This is the distinction people skip, and skipping it is why the wrong technique gets applied.

Type of shiftWhat changesWhat stays the sameConcrete exampleUsual fix
Covariate shiftP(X)P(X) — the inputsP(Y∣X)P(Y|X) — the labelling ruleLine B's brighter, cooler lighting. A scratch is still a scratch.Feature alignment: BatchNorm adaptation, adversarial alignment, MMD
Label (prior) shiftP(Y)P(Y) — class frequenciesP(X∣Y)P(X|Y) — what each class looks likeLine A runs 2% defect rate; Line B runs 11% because it handles rougher stock.Re-weight the loss, or recalibrate the decision threshold
Concept shiftP(Y∣X)P(Y|X) — the labelling rule itselfOften P(X)P(X) roughly holdsLine B's QA team calls a 0.3 mm pit a defect; Line A's team tolerated it.Nothing unsupervised can fix this. You need target labels.
Conditional shiftP(X∣Y)P(X|Y) — class appearanceP(Y)P(Y) — frequenciesDefects on Line B's thinner stock look physically different.Class-conditional alignment; harder, usually needs some labels

Unsupervised domain adaptation only works under covariate shift. If the labelling rule itself has changed, no amount of clever feature alignment will recover it — you are trying to learn a function from data that contains no evidence of it.

Before spending a week on adversarial training, work out which row you are in. The Line B problem is covariate shift: a scratch is still a scratch, it just photographs differently. That is the fixable case.

Measuring how big the shift is

You cannot compute P(X)P(X) directly in 150,000 dimensions, but you can measure the shift with a trick: train a classifier to tell source from target. Label every source image 0 and every target image 1, train a small model on the backbone's features, and look at how well it does on held-out data.

Python
import numpy as npfrom sklearn.linear_model import LogisticRegressionfrom sklearn.model_selection import cross_val_score# feats_s, feats_t: (N, 2048) arrays of backbone featuresX = np.vstack([feats_s, feats_t])d = np.hstack([np.zeros(len(feats_s)), np.ones(len(feats_t))])err = 1 - cross_val_score(LogisticRegression(max_iter=1000),                          X, d, cv=5, scoring="accuracy").mean()proxy_a_distance = 2 * (1 - 2 * err)print(f"domain classifier error {err:.3f}  ->  dA = {proxy_a_distance:.2f}")

The proxy A-distance is dA=2(1−2ϵ)d_A = 2(1 - 2\epsilon), where ϵ\epsilon is that classifier's error rate. Work through the two extremes. If the domains are indistinguishable the classifier is at chance, ϵ=0.5\epsilon = 0.5, and dA=2(1−1)=0d_A = 2(1 - 1) = 0. If they are trivially separable, ϵ=0.02\epsilon = 0.02, and dA=2(1−0.04)=1.92d_A = 2(1 - 0.04) = 1.92 — close to the maximum of 2.

On the Line A / Line B features, the domain classifier scored 97% accuracy, so ϵ=0.03\epsilon = 0.03 and dA=2(1−0.06)=1.88d_A = 2(1 - 0.06) = 1.88. A linear model can separate the two lines almost perfectly from backbone features alone. That number is both the diagnosis and the target: successful adaptation should drive it down.

Four scenarios, distinguished by what labels you have

ScenarioTarget labels availableTypical approachRealistic outcome
Supervised DAPlenty (thousands)Just fine-tune on target data, optionally mixing in sourceBest results. Barely counts as adaptation.
Semi-supervised DAA handful (5–50 per class)Fine-tune on the few labels + pseudo-labelling or alignment on the restRecovers most of the gap. Best value for annotation budget.
Unsupervised DANone, but unlabelled target images are plentifulBatchNorm adaptation, adversarial alignment, MMD, self-trainingTypically recovers 40–70% of the gap.
Domain generalisationNone, and no target data at all at training timeTrain across many source domains; heavy augmentation; domain-invariant lossesWeakest, but the only option when the target is unknown in advance.

The row worth pausing on is semi-supervised. Going from zero target labels to just five per class often buys more than the most sophisticated unsupervised method. For the Line B problem, 31 defect classes × 5 labelled examples = 155 annotations, roughly a day of one metallurgist's time. Before committing to adversarial training, always price out that day.

Batch normalisation adaptation: the cheapest fix that works

Start here, always, because it costs one forward pass and no training at all.

Why frozen BatchNorm breaks under domain shift

A BatchNorm layer stores running estimates of the mean and variance of its inputs, accumulated during training, and at inference it normalises using those stored numbers:

x^=γ⋅x−μrunningσrunning2+ϵ+β\hat{x} = \gamma \cdot \frac{x - \mu_{running}}{\sqrt{\sigma^2_{running} + \epsilon}} + \beta

Those running statistics are a compressed description of the source distribution. Feed target data through and the normalisation is simply wrong.

Take a real measurement from one early-layer channel on the defect detector. On Line A images that channel's activations had μs=0.42\mu_s = 0.42, σs=0.19\sigma_s = 0.19. On Line B images the same channel produced μt=0.61\mu_t = 0.61, σt=0.31\sigma_t = 0.31. Normalising a typical Line B activation of 0.61 with the stored Line A statistics gives:

0.61−0.420.19=1.00\frac{0.61 - 0.42}{0.19} = 1.00

The layer's own statistics say this activation should normalise to 0 — it is exactly the target-domain mean. Instead it comes out at 1.0, a full standard deviation off centre. Every downstream layer now receives systematically shifted inputs, and in a 50-layer network those shifts compound.

The fix: recompute the statistics on target data

Adaptive BatchNorm (AdaBN) simply replaces μrunning,σrunning\mu_{running}, \sigma_{running} with statistics computed on unlabelled target data. No labels, no gradients, no optimiser.

Python
import torch, torch.nn as nndef adapt_batchnorm(model, target_loader, num_batches=50, device="cuda"):    """Recompute BN running statistics on unlabelled target data."""    # Reset the stored statistics and put BN into training mode so    # forward passes update them. Everything else stays in eval mode.    model.eval()                   # dropout and everything else: inference mode    for m in model.modules():        if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d)):            m.reset_running_stats()            m.momentum = None      # None -> cumulative moving average            m.train()    with torch.no_grad():        for i, (images, _) in enumerate(target_loader):   # labels unused            if i >= num_batches:                break            model(images.to(device))    model.eval()    return model

Setting momentum = None makes PyTorch accumulate an exact cumulative average rather than an exponential one, which is what you want for a one-off recalibration. Fifty batches of 64 images is 3,200 unlabelled target images — for stable estimates you want at least a few hundred, and returns flatten past a few thousand.

Line B configurationF1Recovered gapCost
Source model, no adaptation68.5%——
+ AdaBN on 3,200 unlabelled images79.1%41%~20 seconds, no labels
+ AdaBN + adversarial alignment86.4%70%~4 hours training, no labels
Fine-tuned on 155 target labels88.2%77%1 day of annotation
Oracle: fine-tuned on 12,000 target labels94.0%100%4 months of annotation

Recomputing BatchNorm statistics on unlabelled target data recovered 41% of a 25-point gap in twenty seconds. Try it before anything else, every time.

Domain adversarial training: make the features unable to tell the domains apart

AdaBN fixes the normalisation but leaves the features themselves source-specific. Domain adversarial training goes further: it changes what the network computes, so that the representation carries class information but no domain information.

The core idea

Three components share one backbone:

  • A feature extractor GfG_f that maps images to a representation.
  • A label predictor GyG_y that classifies the representation. Trained on source data only, because only source data has labels.
  • A domain discriminator GdG_d that tries to say whether a representation came from source or target. Trained on both, using the free domain label.

The discriminator minimises its own domain-classification loss. The feature extractor maximises it — it is rewarded for producing features the discriminator cannot classify. At equilibrium the features are domain-invariant, and since the label predictor is simultaneously being trained to work on those features, they are also class-discriminative.

The objective is one saddle point:

min⁡Gf,Gymax⁡Gd  Ly(Gy(Gf(xs)),ys)−λ Ld(Gd(Gf(x)),d)\min_{G_f, G_y} \max_{G_d} \; \mathcal{L}_y(G_y(G_f(x_s)), y_s) - \lambda \, \mathcal{L}_d(G_d(G_f(x)), d)

The gradient reversal layer

The elegant implementation trick is a layer that is the identity going forwards and multiplies the gradient by −λ-\lambda going backwards. That single sign flip turns a minimisation into the required maximisation, and lets you train the whole thing with one ordinary optimiser and one backward pass.

Python
import torch, torch.nn as nnfrom torch.autograd import Functionclass GradientReversal(Function):    @staticmethod    def forward(ctx, x, lambd):        ctx.lambd = lambd        return x.view_as(x)              # identity forwards    @staticmethod    def backward(ctx, grad_output):        return -ctx.lambd * grad_output, None   # negate backwardsdef grad_reverse(x, lambd=1.0):    return GradientReversal.apply(x, lambd)class DANN(nn.Module):    def __init__(self, backbone, feat_dim, num_classes):        super().__init__()        self.features = backbone                 # e.g. ResNet-50 body        self.classifier = nn.Linear(feat_dim, num_classes)        self.discriminator = nn.Sequential(            nn.Linear(feat_dim, 1024), nn.ReLU(),            nn.Dropout(0.5),            nn.Linear(1024, 2),                  # source vs target        )    def forward(self, x, lambd=1.0):        f = self.features(x).flatten(1)        return self.classifier(f), self.discriminator(grad_reverse(f, lambd))

The training step feeds source and target batches together:

Python
for step, ((xs, ys), (xt, _)) in enumerate(zip(src_loader, tgt_loader)):    p = step / total_steps    lambd = 2.0 / (1.0 + np.exp(-10.0 * p)) - 1.0    # ramp 0 -> 1    cls_s, dom_s = model(xs.cuda(), lambd)    _,     dom_t = model(xt.cuda(), lambd)    loss_cls = F.cross_entropy(cls_s, ys.cuda())    loss_dom = (F.cross_entropy(dom_s, torch.zeros(len(xs), dtype=torch.long).cuda())              + F.cross_entropy(dom_t, torch.ones(len(xt),  dtype=torch.long).cuda()))    (loss_cls + loss_dom).backward()    optimiser.step(); optimiser.zero_grad()

The λ\lambda schedule matters more than people expect. At p=0p = 0 it gives λ=2/(1+e0)−1=0\lambda = 2/(1+e^{0}) - 1 = 0; at p=0.5p = 0.5, λ=2/(1+e−5)−1=0.987\lambda = 2/(1 + e^{-5}) - 1 = 0.987; at p=1p = 1, λ=0.9999\lambda = 0.9999. So the adversarial pressure starts at exactly zero and ramps smoothly to full strength. Turn it on at full strength from step 1 and the features are pushed to be domain-invariant before they encode anything useful about the classes — the network happily collapses to a constant representation, which is perfectly domain-invariant and perfectly useless.

What DANN gives youWhat it costs you
Needs no target labels at allAdversarial training is unstable; results vary several points across seeds
Architecture-agnostic — bolt it onto any backboneAdds real hyperparameters: λ\lambda schedule, discriminator size, learning-rate balance
One optimiser, one backward passAligns the marginal P(X)P(X) only; can align a target "cat" onto a source "dog" cluster
Consistently strong on standard benchmarksFails badly under label shift, since matching marginals then requires misclassifying

That third limitation deserves emphasis because it is the named failure mode of adversarial adaptation. The discriminator only sees the overall distribution of features; it has no idea which class anything is. A solution that maps target images of class 3 onto the source cluster for class 7 satisfies the discriminator completely. Class-conditional variants exist that condition the discriminator on the predicted label, precisely to close this hole.

Self-training with pseudo-labels

A different angle, and often a stronger one: use the model's own confident predictions on target data as if they were ground truth.

Python
@torch.no_grad()def make_pseudo_labels(model, target_loader, threshold=0.95):    model.eval()    keep_x, keep_y, total = [], [], 0    for images, _ in target_loader:        probs = model(images.cuda()).softmax(dim=1)        conf, pred = probs.max(dim=1)        mask = conf >= threshold        total += len(images)        keep_x.append(images[mask.cpu()]); keep_y.append(pred[mask].cpu())    x, y = torch.cat(keep_x), torch.cat(keep_y)    print(f"kept {len(y)}/{total} = {100*len(y)/total:.1f}% above {threshold}")    return x, y

Run it on the 5,000 unlabelled Line B images with a threshold of 0.95 and 1,240 pass — 24.8%. Spot-checking those 1,240 against a small audited sample shows 91% are correct. So you have just added 1,128 correct labels and 112 wrong ones to your training set.

Whether that is a good trade depends entirely on the errors being random rather than systematic, and they never are. The model is confidently wrong on exactly the cases it misunderstands — and now it trains on those mistakes, becomes more confident in them, and passes them through the threshold more easily next round. This is confirmation bias, and it is the reason naive self-training sometimes gets worse with every iteration.

GuardWhat it prevents
High threshold (0.95–0.99), lowered gradually across roundsAdmitting low-quality labels early, when the model is weakest
Class-balanced selection — take the top-k per class, not the top-k overallThe rich-get-richer collapse where one easy class swallows the pseudo-label set
Regenerate pseudo-labels every round rather than accumulating themLocking in early mistakes permanently
Keep source data in every batch alongside pseudo-labelled target dataDrifting away from the labels you actually trust
Consistency regularisation: require the same prediction under strong augmentationConfident-but-fragile predictions passing the threshold

Pseudo-labelling amplifies whatever the model already believes. Used with a class-balanced, regenerated, high-threshold selection it is one of the strongest adaptation methods available; used naively it is a machine for manufacturing confident errors.

Maximum Mean Discrepancy: align the distributions explicitly

Adversarial training aligns distributions implicitly, through a game. MMD does it directly, with a differentiable distance you can add to the loss — and with no adversary, so no instability.

MMD maps both sets of features into a reproducing kernel Hilbert space and measures the distance between their means there. The empirical squared form, for nn source and mm target samples, is:

MMD2=1n2∑i,jk(xis,xjs)+1m2∑i,jk(xit,xjt)−2nm∑i,jk(xis,xjt)\text{MMD}^2 = \frac{1}{n^2}\sum_{i,j} k(x_i^s, x_j^s) + \frac{1}{m^2}\sum_{i,j} k(x_i^t, x_j^t) - \frac{2}{nm}\sum_{i,j} k(x_i^s, x_j^t)

Read it as: source-to-source similarity, plus target-to-target similarity, minus twice the cross similarity. If the two sets are drawn from the same distribution, the cross term is as large as the within terms and the whole expression goes to zero. With a characteristic kernel such as the Gaussian, MMD=0\text{MMD} = 0 if and only if the distributions are identical.

Python
def gaussian_mmd(source, target, sigmas=(1, 2, 4, 8, 16)):    """Multi-kernel MMD^2 between two (B, D) feature batches."""    x = torch.cat([source, target], dim=0)    d2 = torch.cdist(x, x) ** 2                      # squared distances    # Median heuristic: scale bandwidths by the median pairwise distance.    base = d2.detach().median().clamp(min=1e-8)    k = sum(torch.exp(-d2 / (base * s)) for s in sigmas) / len(sigmas)    n = source.size(0)    k_ss = k[:n, :n].mean()    k_tt = k[n:, n:].mean()    k_st = k[:n, n:].mean()    return k_ss + k_tt - 2 * k_st# Used as an auxiliary loss on the penultimate features:loss = F.cross_entropy(logits_s, y_s) + 0.5 * gaussian_mmd(feat_s, feat_t)

Using several bandwidths at once matters. A single Gaussian kernel is only sensitive to differences at roughly its own length scale; a mixture over five bandwidths spanning a 16× range detects both fine and coarse mismatches. The median heuristic sets the centre of that range from the data rather than from a guess.

Adversarial (DANN)MMD-based
StabilityCan oscillate or collapse; seed-sensitiveStable — it is just an extra loss term
Hyperparametersλ\lambda schedule, discriminator architecture, LR balanceLoss weight and kernel bandwidths
ComputeExtra network, extra forward passO(B2)O(B^2) kernel matrix per batch — cheap for B≤128B \le 128
Alignment strengthMatches full distributions in principleMatches kernel mean embeddings; weaker but more predictable
Typical useWhen you can afford to tune itWhen you want a reliable improvement with little tuning

Putting it together on a real benchmark

Office-31 is the standard test bed: 31 object categories photographed in three domains — Amazon product shots on white backgrounds (2,817 images), a DSLR in an office (498 images), and a low-resolution webcam in the same office (795 images). The hard transfer is Amazon → Webcam, because studio product photography and a grainy webcam are genuinely different worlds.

Python
for step in range(total_steps):    xs, ys = next(src_iter)    xt, _  = next(tgt_iter)          # target labels never touched    p = step / total_steps    lambd = 2.0 / (1.0 + np.exp(-10.0 * p)) - 1.0    fs = model.features(xs.cuda()).flatten(1)    ft = model.features(xt.cuda()).flatten(1)    loss = F.cross_entropy(model.classifier(fs), ys.cuda())      # supervised    loss = loss + 0.5 * gaussian_mmd(fs, ft)                     # alignment    loss = loss + lambd * domain_adversarial_loss(fs, ft)        # adversarial    loss.backward(); optimiser.step(); optimiser.zero_grad()

Published results for single methods on this pair, each with a ResNet-50 backbone and no target labels at all (as reported by Long et al., Conditional Adversarial Domain Adaptation, 2018):

Method (Amazon → Webcam)Target accuracyTarget labels used
Source-only ResNet-5068.4%0
DAN (multi-kernel MMD alignment)80.5%0
DANN (domain adversarial)82.0%0
JAN (joint MMD over features and predictions)85.4%0
CDAN+E (class-conditional adversarial)94.1%0

Two lessons sit in that table. Plain distribution alignment, MMD or adversarial, recovers 12 to 14 of the points lost to the domain gap without a single target label. And the biggest jump comes from conditioning the alignment on the predicted class — exactly the fix for the "target cat onto source dog" failure described above. Treat these as benchmark numbers from careful research code; on your own data, expect smaller and noisier gains.

How to run this on a real deployment

The practical consequence is a sequence, ordered by cost, with a stopping rule at each step.

  1. Diagnose the shift type before touching anything. Compare class frequencies between source and target if you have any target labels at all — a big difference means label shift, and threshold recalibration will beat every method in this topic. Sample fifty target images and check by eye whether the labelling rule still holds; if it has changed, that is concept shift and unsupervised adaptation is provably hopeless.
  2. Measure dAd_A. A proxy A-distance below about 0.5 means the shift is mild and your problem is probably something else. Above 1.5 means the domains are trivially separable and adaptation has real work to do. Track this number as you adapt; it should fall.
  3. Run AdaBN. Twenty seconds, no labels, no risk. On the defect detector it recovered 41% of the gap. There is no reason not to do it.
  4. Add an MMD term before you try anything adversarial. It is stable, has two hyperparameters instead of five, and on Office-31 Amazon → Webcam it came within 1.5 points of DANN. Reach for adversarial training when MMD has plateaued and you have budget to tune it.
  5. Then pseudo-label, carefully. Class-balanced top-k selection, labels regenerated each round, source data still in every batch. Like class-conditional alignment, it brings back the class-level information that plain distribution matching ignores.
  6. Always price the annotation alternative. Three to five labels per class is often one afternoon of work, and it is frequently competitive with weeks of tuning an unsupervised method. If someone can label, let them.

One warning that applies to every step: you cannot validate any of this on a labelled target validation set, because if you had one you would be doing supervised adaptation. Practitioners routinely tune unsupervised methods on target test labels and report the result as unsupervised, which is meaningless. Use whatever you legitimately have — source validation accuracy, the domain classifier's dAd_A, the fraction of target samples above the confidence threshold, prediction entropy on target data — and accept the honest cost of that constraint: unsupervised adaptation is much harder to tune reliably than the benchmark tables make it look.