Data augmentation in practice
Be able to choose augmentations that match the domain and measure the effect.
Prerequisites
Intuition
Augmentation creates more training examples by changing existing ones in ways that do not change the label. It is the most effective regularisation there is when data is limited.
The decisive question for every augmentation: does it preserve the label?
| Augmentation | Cat/dog | Digits (MNIST) | Medical images |
|---|---|---|---|
| Horizontal flip | yes | no (a 2 does not stay a 2) | often no (left/right organ) |
| Rotation ±15° | yes | yes | yes |
| Rotation 180° | yes | no (a 6 becomes a 9) | no |
| Colour jitter | yes | irrelevant | no (colour is diagnostic) |
| Cropping | yes | partly | no (it can cut away the lesion) |
The most common slip is copying an augmentation recipe from ImageNet to a domain where it does not hold. A horizontal flip is standard in image classification and directly harmful in digit recognition.
Formal
Augmentations that mix examples — a category of their own, and the strongest ones:
| Method | How | Effect |
|---|---|---|
| Mixup | , the same on the labels | smoother decision boundaries, better calibration |
| CutMix | cut a rectangle out of one image, paste it from another; the labels are weighted by area | keeps the local detail |
| RandAugment | pick random operations with strength | two parameters instead of twenty |
| TrivialAugment | one random operation at a random strength | zero parameters, often just as good |
The last row is worth noting: TrivialAugment showed that the whole search for augmentation policies could largely be replaced by pure chance, with comparable results.
Where in the pipeline:
| Step | Augmentation |
|---|---|
| Training | yes — at random per epoch |
| Validation | no — otherwise you are measuring on distorted data |
| Test | no, except for deliberate TTA |
Test-time augmentation (TTA) is the exception: augment the test image a few times, run the model on each and average. It typically gives 0.5–2 percentage points for an increased inference cost.
Measure, do not assume. The effect varies enormously with the domain and the amount of data:
| Amount of data | The typical effect of heavy augmentation |
|---|---|
| Hundreds of images | very large — often decisive |
| Tens of thousands | moderate |
| Millions | small, sometimes negative |
Other modalities:
| Modality | What works |
|---|---|
| Text | synonym swaps, back-translation, word deletion — carefully, word order carries meaning |
| Audio | noise, time stretching, pitch, SpecAugment |
| Time series | jitter, scaling, window shifting — never a flip in time |
| Tabular | SMOTE for the minority class, noise on the numeric columns |
The row about time series matters: flipping in time reverses cause and effect.
The check that exposes bad augmentations: look at 20 augmented examples with your own eyes. Can you still tell the right label? If you cannot, neither can the model.
Code
import torch, torch.nn as nn
from torchvision.transforms import v2
# Images: training and evaluation are DIFFERENT pipelines
train_tf = v2.Compose([
v2.RandomResizedCrop(224, scale=(0.6, 1.0)),
v2.RandomHorizontalFlip(), # ONLY if flipping preserves the label
v2.TrivialAugmentWide(), # zero parameters, a strong baseline
v2.ToDtype(torch.float32, scale=True),
v2.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
v2.RandomErasing(p=0.25),
])
eval_tf = v2.Compose([
v2.Resize(256), v2.CenterCrop(224), # deterministic — no randomness
v2.ToDtype(torch.float32, scale=True),
v2.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
])
# Mixup: mix both the images and the labels
def mixup(x, y, n_classes, alpha=0.2):
lam = float(torch.distributions.Beta(alpha, alpha).sample())
perm = torch.randperm(x.size(0), device=x.device)
xm = lam * x + (1 - lam) * x[perm]
y1 = nn.functional.one_hot(y, n_classes).float()
return xm, lam * y1 + (1 - lam) * y1[perm]
def soft_cross_entropy(logits, soft_targets):
return -(soft_targets * torch.log_softmax(logits, dim=-1)).sum(dim=-1).mean()
# Test-time augmentation
def tta(model, image, transforms, n=5):
model.eval()
with torch.no_grad():
p = torch.stack([torch.softmax(model(t(image).unsqueeze(0)), -1)
for t in transforms[:n]])
return p.mean(0)
# MEASURE the effect — do not assume it
def compare_augmentation(train_and_evaluate, recipes: dict, seeds=(0, 1, 2)):
import numpy as np
for name, tf in recipes.items():
v = [train_and_evaluate(tf, seed=s) for s in seeds]
print(f" {name:<18} {np.mean(v):.4f} ± {np.std(v):.4f}")
# Look at 20 augmented examples with your own eyes
def inspect(dataset, tf, n=20, out="inspect.png"):
from torchvision.utils import make_grid, save_image
images = torch.stack([tf(dataset[i][0]) for i in range(n)])
save_image(make_grid(images, nrow=5, normalize=True), out)
print(f"saved {out} — can YOU still see the right label?")
inspect is the cheapest quality control there is. An augmentation that makes the image unintelligible to you makes it unintelligible to the model too, and then you are training on noise.
Mastery means
- Chooses augmentations that preserve the label
- Measures the effect instead of assuming it
- Knows which ones belong in training and which in evaluation
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — mixup: Beyond Empirical Risk Minimization — arXiv (open access; licence per article)
- arXiv — TrivialAugment: Tuning-free Yet State-of-the-Art Data Augmentation — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause