Transfer Learning and Pretraining

Handling Small Datasets Effectively


A researcher has 240 labelled histopathology tiles across four tissue classes. It took a pathologist three weeks to annotate them, and there will be no more. She splits 80/20, fine-tunes a ResNet-50, and reports 91% validation accuracy. Then she tries a new augmentation recipe and gets 94%. Three points better — worth writing up.

Her supervisor asks her to re-run both with a different random seed. The original recipe now scores 94%. The new recipe scores 88%.

Here is the arithmetic she should have done first. The validation set contains 48 images, so a single image is worth 1/48=2.081/48 = 2.08 percentage points. The standard error on an accuracy estimate of 0.91 from 48 samples is

0.91×0.0948=0.001706=0.0413\sqrt{\frac{0.91 \times 0.09}{48}} = \sqrt{0.001706} = 0.0413

which gives a 95% confidence interval of roughly ±8.1 points: 82.9% to 99.1%. Her three-point improvement sat entirely inside the noise. She had not measured an improvement; she had measured which images happened to land in the validation split.

Small datasets are hard in two distinct ways, and this is the second one. Everyone knows that training on little data overfits. Far fewer people notice that evaluating on little data is almost meaningless, which means you cannot even tell whether your fixes are working.

Six levers when there will be no more data240 labelled tilesAugment harderas data shrinksMixup: train on blendsCutMix: paste in a rectangleDropout and weight decayEarly stoppingon a real splitK-fold for anumber you trust
On a 48-image validation set, 91% and 94% differ by one and a half images — cross-validation is what tells them apart.

Why small datasets break models

The capacity mismatch

ResNet-50 has 25.6 million trainable parameters. With 240 training images that is 106,667 free parameters per example. A system with a hundred thousand degrees of freedom per constraint is not solving a problem; it is picking arbitrarily from an enormous space of functions that all fit the training data perfectly. Almost every one of them is memorisation.

The tell is a widening gap. Training accuracy climbs to 99%+ while validation stalls or falls. The model is not learning that nuclei arranged in glandular structures indicate one class; it is learning that image 47 is class 2.

Overfitting is not something that happens if you train too long. It is the default outcome whenever capacity exceeds the constraints the data supplies — training longer just lets the model finish the job.

Transfer learning helps, but it does not finish the job

Pretraining collapses the effective problem. Freeze the backbone and train a head from 2048 features to 4 classes and you have 2048×4+4=8,1962048 \times 4 + 4 = 8{,}196 trainable parameters, or 34 per training image. That is a completely different regime, and it is why transfer learning turns a hopeless problem into a workable one.

But three things remain broken, and each needs its own fix:

  • Domain gap. ImageNet features describe photographs of everyday objects. Histopathology is textural, has no canonical orientation, and its diagnostic signal lives in nuclear morphology at a scale ImageNet never had to represent. Frozen features are a good starting point and a poor final answer.
  • Fine-tuning reopens the capacity problem. The moment you unfreeze layer4 you are back to 15 million parameters on 240 images.
  • Evaluation stays unreliable. Pretraining does nothing whatsoever for the width of your confidence interval.

Augmentation: manufacturing more signal from the same images

Augmentation applies label-preserving transformations so the model sees a different version of each image every epoch. It attacks the capacity mismatch from the data side rather than the model side.

Strength should scale inversely with dataset size

Images per classGeometricPhotometricAdvancedDropout
Under 50RandomResizedCrop(0.5–1.0), flips, rotate ±20°Strong jitter (0.4)Mixup + CutMix + RandomErasing0.5
50–200RandomResizedCrop(0.6–1.0), flips, rotate ±15°Moderate jitter (0.3)RandAugment, light Mixup0.4
200–1,000RandomResizedCrop(0.7–1.0), flipsLight jitter (0.2)RandomErasing0.3
Over 1,000RandomResizedCrop(0.8–1.0), flipsMinimalOptional0.1–0.2

For histopathology specifically, tiles have no canonical orientation — a slide can be placed on the microscope stage any way round — so full 360° rotation and both flips are valid, giving eight free symmetries. But colour jitter must be handled carefully, because haematoxylin and eosin stain intensity carries diagnostic information. This is the general rule in disguised form:

Every augmentation is a claim that the label is invariant to that transformation. When the claim is false, you have not added regularisation — you have injected label noise on purpose.

A flip is free on tissue tiles and destructive on chest X-rays, where mirroring moves the heart to the wrong side. Rotation past 20° is fine on satellite imagery and wrong on handwritten digits, where a rotated 6 becomes a 9. Check every transform against a domain expert's judgement before enabling it.

Mixup: train on blends of two images

Mixup takes two training examples and interpolates both the inputs and the labels:

x~=λxi+(1−λ)xj,y~=λyi+(1−λ)yj,λ∼Beta(α,α)\tilde{x} = \lambda x_i + (1-\lambda) x_j, \qquad \tilde{y} = \lambda y_i + (1-\lambda) y_j, \qquad \lambda \sim \text{Beta}(\alpha, \alpha)

Python
import numpy as np, torch, torch.nn.functional as Fdef mixup_batch(x, y, alpha=0.2):    lam = np.random.beta(alpha, alpha)    idx = torch.randperm(x.size(0), device=x.device)    return lam * x + (1 - lam) * x[idx], y, y[idx], lamdef mixup_loss(logits, y_a, y_b, lam):    return lam * F.cross_entropy(logits, y_a) + \           (1 - lam) * F.cross_entropy(logits, y_b)# in the training loopmixed, y_a, y_b, lam = mixup_batch(images, labels, alpha=0.2)loss = mixup_loss(model(mixed), y_a, y_b, lam)

The detail almost everyone gets wrong is α\alpha. With α=0.2\alpha = 0.2 the Beta distribution is U-shaped — most draws land near 0 or near 1, so most "mixed" images are 95% one image with a faint ghost of another. That is deliberate: it applies gentle pressure without destroying the signal. Setting α=1.0\alpha = 1.0 makes Beta uniform, so 50/50 blends become common, and on a small fine-grained dataset that usually hurts. Start at 0.2 and only raise it if the train–validation gap stays wide.

Mixup also forces linear behaviour between training points, which is why it improves calibration — models trained with it are much less prone to being 99.9% confident about nonsense.

CutMix: paste a rectangle from one image into another

CutMix cuts a random rectangle out of image A and pastes in the corresponding region from image B, mixing labels in proportion to the pasted area.

Python
def cutmix_batch(x, y, alpha=1.0):    lam = np.random.beta(alpha, alpha)    idx = torch.randperm(x.size(0), device=x.device)    _, _, H, W = x.shape    cut_ratio = np.sqrt(1.0 - lam)    cw, ch = int(W * cut_ratio), int(H * cut_ratio)    cx, cy = np.random.randint(W), np.random.randint(H)    x1, x2 = np.clip(cx - cw // 2, 0, W), np.clip(cx + cw // 2, 0, W)    y1, y2 = np.clip(cy - ch // 2, 0, H), np.clip(cy + ch // 2, 0, H)    x[:, :, y1:y2, x1:x2] = x[idx, :, y1:y2, x1:x2]    # Recompute lambda from the ACTUAL pasted area, not the sampled one.    lam = 1.0 - ((x2 - x1) * (y2 - y1) / (W * H))    return x, y, y[idx], lam

That recomputation is the step people omit, and here is why it matters. Sample λ=0.6\lambda = 0.6 on a 224×224 image. Then 1−0.6=0.632\sqrt{1 - 0.6} = 0.632, so the cut box is 224×0.632=141224 \times 0.632 = 141 pixels square, an area of 1412=19,881141^2 = 19{,}881 out of 2242=50,176224^2 = 50{,}176 — a fraction of 0.396. The true mixing ratio is therefore 1−0.396=0.6041 - 0.396 = 0.604, not 0.600. Clipping at image boundaries makes the discrepancy much larger: a box centred near a corner may have half its area cut off, so the label would be badly wrong if you used the sampled value.

CutMix tends to beat Mixup where the object of interest is localised, because the pasted patch is a genuine piece of a real image rather than a translucent overlay. Mixup tends to win on textural data. Using both, chosen at random per batch, is a common and effective default.

Text augmentation

Text has no equivalent of a rotation, because small edits change meaning easily. Four techniques, in rough order of safety:

TechniqueHow it worksRisk
Back-translationTranslate to another language and back: "the film was dull" → German → "the movie was boring"Safest. Preserves meaning, genuinely varies phrasing. Needs a translation model.
Contextual word replacementMask a word and let a language model propose a replacement that fits the contextLow. Can flip sentiment — check that "good" does not become "bad".
Synonym replacementSwap words for thesaurus synonymsModerate. Ignores context; "bank" becomes "riverbank" in a finance sentence.
Random deletion / swapDelete or reorder ~10% of tokensHighest. Destroys negation and syntax: dropping "not" inverts the label.

The negation trap is worth stating plainly: on sentiment data, random deletion that removes "not" from "this was not good" produces a positive-looking sentence still labelled negative. On a 240-example dataset a handful of those is enough to matter.

Regularisation: limiting what the model can memorise

Dropout

Dropout zeroes a random fraction of activations each forward pass, so no single neuron can be relied upon and the network is forced to spread its representation across many units. It is also an implicit ensemble: each training step effectively trains a different thinned network.

Python
import torch.nn as nnmodel.fc = nn.Sequential(    nn.Dropout(0.5),                 # heaviest right after the features    nn.Linear(2048, 256),    nn.BatchNorm1d(256),    nn.ReLU(inplace=True),    nn.Dropout(0.3),                 # lighter in the narrower layer    nn.Linear(256, 4),)

Two rules. Higher rates on wider layers, lower on narrow ones — dropping 50% of a 16-unit layer removes too much information. And dropout must be off at evaluation time; model.eval() handles this, and forgetting it gives you a model whose predictions change between identical calls.

Weight decay

Weight decay penalises large weights, biasing the model towards smoother functions. Prefer AdamW over Adam with weight_decay: in Adam the decay term is folded into the gradient and then divided by the adaptive scaling, so parameters with large gradients get less decay — which is not what you asked for. AdamW applies it as a separate, correct step.

Python
optimiser = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-2)

On 240 images, 1e-2 is a reasonable starting value — an order of magnitude stronger than the 1e-4 you would use on a large dataset. A stronger variant penalises distance from the pretrained weights instead of from zero, which keeps the model near a solution you know is good rather than near the origin, which you have no reason to prefer.

Early stopping, done correctly

Python
best_acc, patience, bad = 0.0, 7, 0for epoch in range(60):    train_one_epoch(model, train_dl, optimiser, criterion)    acc = validate(model, val_dl)    if acc > best_acc + 1e-4:        best_acc, bad = acc, 0        torch.save(model.state_dict(), "best.pt")    else:        bad += 1        if bad >= patience:            print(f"stopping at epoch {epoch}, best {best_acc:.3f}")            breakmodel.load_state_dict(torch.load("best.pt"))     # NOT the last epoch

Restoring the best checkpoint rather than keeping the final one is the part that gets skipped, and skipping it discards the entire benefit. There is also a subtler cost worth naming: choosing the stopping epoch by validation accuracy means your validation score is now optimistically biased, because you selected on it. With 48 validation images that bias is substantial. If a number needs to be trustworthy, keep a separate test set you touch exactly once.

Learning rate scheduling

Python
scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(    optimiser, T_0=10, T_mult=2, eta_min=1e-6)

Cosine annealing decays the rate smoothly from its initial value to near zero, which lets the model take large exploratory steps early and settle precisely later. Warm restarts periodically jump the rate back up, knocking the model out of sharp minima; sharp minima generalise worse, and on small data you are surrounded by them.

Cross-validation: how to get a number you can trust

This is the fix for the researcher's real problem. Instead of one 80/20 split, partition the data into kk folds, train kk models each holding out a different fold, and pool the results. Every example is used for validation exactly once.

The effect on the confidence interval is direct. Her single split evaluated on 48 images gave ±8.1 points. Five-fold cross-validation evaluates on all 240:

0.91×0.09240=0.0185⇒±3.6 points\sqrt{\frac{0.91 \times 0.09}{240}} = 0.0185 \quad \Rightarrow \quad \pm 3.6 \text{ points}

The error bar shrinks by a factor of 5≈2.24\sqrt{5} \approx 2.24, and you get a per-fold standard deviation as well — which is itself diagnostic. If your five folds score 89%, 91%, 90%, 92%, 90%, the model is stable. If they score 78%, 94%, 85%, 96%, 82%, the model is not learning a reliable rule and the mean is hiding that.

Python
import numpy as npfrom sklearn.model_selection import StratifiedKFoldskf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)scores = []for fold, (tr_idx, va_idx) in enumerate(skf.split(paths, labels)):    model = build_model()                      # fresh model every fold    train(model, subset(tr_idx))    acc = validate(model, subset(va_idx))    scores.append(acc)    print(f"fold {fold}: {acc:.4f}")print(f"{np.mean(scores):.4f} +/- {np.std(scores):.4f}")

Stratified folds preserve each class's proportion in every fold. With 240 images split 168/40/22/10 across four classes, a random 5-fold split would give each fold about 2 examples of the rarest class — and by chance some folds would get 0, making that class's recall undefined and the fold's accuracy incomparable. Stratification guarantees each fold gets 2.

Two further warnings. Build a fresh model each fold; reusing one leaks information from previous folds' validation data. And if your data has group structure — several tiles from the same patient, several photos of the same object — you must split by group, not by example, or near-duplicates land on both sides of the split and your score becomes fiction. Use StratifiedGroupKFold for that.

Class imbalance

Her 240 tiles split 168 / 40 / 22 / 10. A model predicting the majority class for everything scores 168/240=70%168/240 = 70\% accuracy while being completely useless. Three tools address this.

Class weights

Weight each class's loss contribution inversely to its frequency: wc=NK⋅ncw_c = \frac{N}{K \cdot n_c} for NN total examples and KK classes.

ClassCountCalculationWeight
0168240 / (4 × 168)0.357
140240 / (4 × 40)1.500
222240 / (4 × 22)2.727
310240 / (4 × 10)6.000
Python
from sklearn.utils.class_weight import compute_class_weightw = compute_class_weight("balanced", classes=np.unique(labels), y=labels)criterion = nn.CrossEntropyLoss(weight=torch.tensor(w, dtype=torch.float).cuda())

One misclassified example of class 3 now costs 16.8 times as much as one of class 0. The normalisation by KK keeps the average weight at 1, so your loss stays on a comparable scale.

Oversampling

Python
from torch.utils.data import WeightedRandomSamplersample_w = [w[label] for label in labels]sampler = WeightedRandomSampler(sample_w, num_samples=len(labels), replacement=True)loader = DataLoader(dataset, batch_size=16, sampler=sampler)   # no shuffle=True

This draws minority examples more often, so batches are roughly balanced. Note that with only 10 examples of class 3, oversampling shows the model the same ten images repeatedly — strong augmentation is essential or it will memorise them exactly. Do not use class weights and oversampling together at full strength; you will double-count the correction and swing into over-predicting the rare class.

Focal loss

Focal loss reshapes the loss so easy examples contribute less: FL(pt)=−(1−pt)γlog⁡(pt)\text{FL}(p_t) = -(1 - p_t)^\gamma \log(p_t).

Work the numbers with γ=2\gamma = 2. An easy example the model already gets right with pt=0.9p_t = 0.9: standard cross-entropy is −ln⁡(0.9)=0.105-\ln(0.9) = 0.105, and focal loss multiplies it by (1−0.9)2=0.01(1-0.9)^2 = 0.01, giving 0.00105. A hard example with pt=0.3p_t = 0.3: cross-entropy is −ln⁡(0.3)=1.204-\ln(0.3) = 1.204, multiplied by (1−0.3)2=0.49(1-0.3)^2 = 0.49, giving 0.590.

Under plain cross-entropy the hard example is worth 11.4 times the easy one. Under focal loss it is worth 0.590/0.00105=5620.590 / 0.00105 = 562 times as much — a 49-fold shift in relative attention towards the examples the model is still getting wrong.

Python
class FocalLoss(nn.Module):    def __init__(self, gamma=2.0, weight=None):        super().__init__()        self.gamma, self.weight = gamma, weight    def forward(self, logits, target):        ce = F.cross_entropy(logits, target, weight=self.weight, reduction="none")        p_t = torch.exp(-ce)                       # probability of true class        return ((1 - p_t) ** self.gamma * ce).mean()

Use focal loss when imbalance is severe (worse than about 1:20) and the rare classes are genuinely hard. When the rare classes are easy but simply infrequent, class weights are simpler and work as well.

Everything together

Python
skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)fold_scores = []for fold, (tr_idx, va_idx) in enumerate(skf.split(paths, labels)):    model = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)    for p in model.parameters():        p.requires_grad = False                        # freeze first    model.fc = nn.Sequential(nn.Dropout(0.5), nn.Linear(2048, 4))    model = model.cuda()    w = compute_class_weight("balanced", classes=np.unique(labels[tr_idx]),                             y=labels[tr_idx])    # The sampler below already balances the classes, so the loss gets no class    # weights -- using both at full strength would double-count the correction.    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)    sampler = WeightedRandomSampler([w[l] for l in labels[tr_idx]],                                    num_samples=len(tr_idx), replacement=True)    train_dl = DataLoader(Subset(ds_train, tr_idx), batch_size=16, sampler=sampler)    val_dl = DataLoader(Subset(ds_eval, va_idx), batch_size=32)    # Phase 1: head only. Establishes the baseline and warms up the head.    opt = torch.optim.AdamW(model.fc.parameters(), lr=1e-3, weight_decay=1e-2)    best = train_with_early_stopping(model, train_dl, val_dl, opt,                                     criterion, epochs=25, patience=7)    # Phase 2: unfreeze the last block at a 10x lower rate, with Mixup.    for p in model.layer4.parameters():        p.requires_grad = True    opt = torch.optim.AdamW(        [{"params": model.fc.parameters(),     "lr": 1e-4},         {"params": model.layer4.parameters(), "lr": 3.8e-5}], weight_decay=1e-2)    best2 = train_with_early_stopping(model, train_dl, val_dl, opt, criterion,                                      epochs=25, patience=7, use_mixup=True)    fold_scores.append(max(best, best2))print(f"CV accuracy {np.mean(fold_scores):.4f} +/- {np.std(fold_scores):.4f}")

Notice that phase 2 is kept only if it beats phase 1 — max(best, best2). On 240 images unfreezing 15 million parameters often loses, and the pipeline should let the data decide rather than assuming.

What this means when you build something

The instinct on a small dataset is to reach for more techniques. The correct instinct is to reach for a better measurement first, because without one you cannot tell which techniques are helping.

  1. Compute your error bar before your first experiment. p(1−p)/n\sqrt{p(1-p)/n} on your validation size, times two. If it is wider than the improvements you expect to make, single-split evaluation is worthless and cross-validation is not optional.
  2. Cross-validate everything, and report the standard deviation. A mean without a spread is not a result. Five folds costs five times the compute, which on 240 images is still minutes.
  3. Fix the split before you fix the model. Stratify by class; group by patient, session or source. Leakage across the split inflates every number you will produce afterwards, and it is invisible unless you look for it.
  4. Start with the frozen backbone. It is the configuration least able to overfit, it takes minutes, and it is the number every later experiment must beat.
  5. Add regularisation in order of cost: augmentation, then dropout and weight decay, then Mixup or CutMix, then early stopping. Add one at a time and keep it only if cross-validated accuracy improves by more than the fold standard deviation.
  6. Handle imbalance explicitly, and stop reporting accuracy. With a 168/40/22/10 split, report per-class recall and macro-F1. Accuracy will look respectable while the model never once predicts the rarest class.
  7. Price the alternative honestly. Going from 240 to 500 labelled examples typically buys more than every technique in this list combined. Before spending three weeks on methodology, work out whether three more days of a pathologist's time is available — it usually wins.

The researcher re-ran her two recipes under 5-fold cross-validation. The original scored 90.4% ± 2.1, the new one 90.8% ± 2.6. There was no improvement — there never had been. What she gained was the ability to tell.