Geometric augmentations
Flips, crops, rotations and shifts multiply a small dataset by teaching the model which changes should not change its answer, and picking the wrong ones teaches it something false.
- 12 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.
Geometric augmentation means showing the model tilted, flipped and cropped copies of each photo, so it stops memorising position.
A mango is a mango upside down, sideways, close up, or half behind somebody's hand. You knew that as a child, because you saw mangoes from every angle for years.
A model has seen four hundred mangoes, all centred, all upright. It learned "mango" and "centred and upright" as one idea. Move the fruit and it hesitates.
Augmentation is showing it the angles you never photographed.
The rule that governs everything
Some changes leave the answer alone. Some destroy it.
a photo of a mango
|
+----+----+-----+------+
| | | | |
tilted flip crop zoom shifted -> still a mango -> SAFE
a photo of the number 6
|
turned upside down
|
it is a 9 -> label is now WRONG -> UNSAFEThat is the entire skill. For each transformation you are about to apply, ask: after this change, is the label still true?
Nobody can answer that for you, because it depends on your task. Mirroring a cat is harmless. Mirroring a road sign with an arrow on it is a lie. Mirroring a chest X-ray moves the heart to the wrong side, which is exactly the thing a radiologist is checking.
Why augmentation works at all
You have eight hundred photos. Every epoch, the model sees each photo again. By the fifth pass it has begun to memorise them, including details unrelated to the task.
With augmentation, it never sees the same picture twice. Photo 41 arrives tilted four degrees and cropped slightly left, then next epoch tilted the other way. The parts that stay constant across all those versions are the parts worth learning.
This is the cheapest way to fight memorisation. It costs no new labels and no new data collection.
The common transformations
- Horizontal flip — mirror left to right. Safe for most objects, unsafe for text, digits, and anything where left and right differ.
- Random resized crop — take a random rectangle and stretch it to full size. Teaches tolerance to zoom and framing.
- Rotation — turn by a few degrees. Small angles for photos taken by a person, any angle for microscope or satellite images.
- Translation — shift the picture. Often the most useful and least talked about.
- Vertical flip — upside down. Safe for satellite and microscope images, rarely safe for photographs of the world.
The honest part
More augmentation is not better. Push it far enough and you train on pictures your model will never meet. That wastes capacity and slows learning.
The developer section runs an experiment where turning on one popular augmentation drops accuracy from perfect to near-guessing. It is one line of code, on by default in many tutorials, and wrong for that dataset.
Remember this
- Augmentation shows altered copies of the same photos, so the model stops memorising them.
- A transformation is allowed only if the label survives it.
- The right set depends on your data, and copying somebody else's list is how the damage happens.
What to learn next
- Mixup and CutMix — augmentations that change the label on purpose, in a controlled way.
- RandAugment and AutoAugment — letting a search pick the transformation list instead of you.
- Data augmentation — the same idea across text, audio and tabular data.
Developer — Code and libraries.
Setup
pip install torch torchvision numpyWritten and run against torch 2.13.0 (CPU), torchvision 0.28.0, numpy 2.2.6. No downloads, no pretrained weights. Both scripts run in seconds.
Everything uses torchvision.transforms.v2, the current API. The v1 transforms still work, but only v2 transforms geometry on masks and boxes alongside the image. See torchvision transforms v2.
Geometry moves the labels too
For classification, a flip changes the pixels and leaves the label alone. For detection and segmentation the label is geometry, so it has to move with the image. That is what tv_tensors are for.
import torch
from torchvision.transforms import v2
from torchvision import tv_tensors
torch.manual_seed(0)
img = tv_tensors.Image(torch.zeros(3, 40, 40, dtype=torch.uint8))
img[:, 8:20, 5:25] = 255 # a bright box in the top-left area
boxes = tv_tensors.BoundingBoxes([[5, 8, 25, 20]], format="XYXY", canvas_size=(40, 40))
flip = v2.RandomHorizontalFlip(p=1.0) # p=1 so the demo is not a coin toss
out_img, out_boxes = flip(img, boxes)
print("original box :", boxes.tolist())
print("flipped box :", out_boxes.tolist())
rot = v2.RandomRotation(degrees=(90, 90)) # a fixed 90 degrees, for a readable result
r_img, r_boxes = rot(img, boxes)
print("rotated box :", [[round(v) for v in b] for b in r_boxes.tolist()])
print("image dtype/shape unchanged:", r_img.dtype, tuple(r_img.shape))original box : [[5, 8, 25, 20]] flipped box : [[15, 8, 35, 20]] rotated box : [[8, 15, 20, 35]] image dtype/shape unchanged: torch.uint8 (3, 40, 40)
The box started at x from 5 to 25 on a 40-wide canvas. After the mirror it sits at x from 15 to 35, which is 40 - 25 to 40 - 5. The transform did the arithmetic because the tensor was wrapped in tv_tensors.BoundingBoxes and carried its format and canvas_size along.
Pass a plain tensor of coordinates instead and it is treated as an image-like array. No error, wrong boxes, and a detector that trains to nothing.
The experiment worth running once
Two classes: an arrow pointing left, and the same arrow pointing right. Nothing else differs.
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.data import TensorDataset, DataLoader
from torchvision.transforms import v2
rng = np.random.default_rng(0)
def arrow(points_right):
img = np.zeros((32, 32), "float32")
row = int(rng.integers(10, 22))
img[row-1:row+2, 6:26] = 1.0 # the shaft
for i in range(7): # the head, at one end only
if points_right: img[row-i:row+i+1, 25-i] = 1.0
else: img[row-i:row+i+1, 6+i] = 1.0
return img
def make(n):
lab = rng.integers(0, 2, n)
x = torch.tensor(np.stack([arrow(bool(c)) for c in lab])).unsqueeze(1)
return x, torch.tensor(lab)
Xtr, ytr = make(800); Xte, yte = make(400)
def train(aug, seed=0):
torch.manual_seed(seed)
net = nn.Sequential(nn.Conv2d(1,16,3,padding=1), nn.ReLU(), nn.MaxPool2d(2),
nn.Conv2d(16,32,3,padding=1), nn.ReLU(), nn.MaxPool2d(2),
nn.Flatten(), nn.Linear(32*8*8, 2))
opt = torch.optim.AdamW(net.parameters(), lr=1e-3)
dl = DataLoader(TensorDataset(Xtr, ytr), batch_size=32, shuffle=True)
for _ in range(6):
net.train()
for xb, yb in dl:
if aug is not None: xb = aug(xb)
opt.zero_grad(); F.cross_entropy(net(xb), yb).backward(); opt.step()
net.eval()
with torch.no_grad(): return (net(Xte).argmax(1) == yte).float().mean().item()
print("no augmentation :", round(train(None), 3))
print("random horizontal flip :", round(train(v2.RandomHorizontalFlip(p=0.5)), 3))
print("small rotations (+/- 12) :", round(train(v2.RandomRotation(degrees=12)), 3))no augmentation : 1.0 random horizontal flip : 0.645 small rotations (+/- 12) : 1.0
Deterministic with the seeds set.
Reading that
RandomHorizontalFlip(p=0.5) took a perfect model down to 0.645. Chance is 0.5. Half of every batch arrived mirrored, still carrying its original label. The model was told, repeatedly, that a left arrow is a right arrow.
Rotation left it at 1.0. A twelve-degree tilt does not turn a left arrow into a right one. The label survived and the augmentation was free.
Neither transformation is good or bad in itself. The dataset decides. On photographs of animals the flip would have helped and cost nothing.
This failure is quiet in real projects. Your dataset is not two arrows. The damage is a few points of accuracy, never traced back to the copied transform list.
A defensible starting pipeline
train_tf = v2.Compose([
v2.ToImage(), # PIL or ndarray -> tensor image
v2.RandomResizedCrop(224, scale=(0.7, 1.0), antialias=True),
v2.RandomHorizontalFlip(p=0.5), # DELETE THIS LINE if left/right matters
v2.ToDtype(torch.float32, scale=True),
v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
eval_tf = v2.Compose([ # no randomness at evaluation time
v2.ToImage(),
v2.Resize(256, antialias=True),
v2.CenterCrop(224),
v2.ToDtype(torch.float32, scale=True),
v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])Order matters. Geometry first while the tensor is still uint8 and cheap to move, dtype conversion next, normalisation last. v2.ToTensor is deprecated in favour of the ToImage plus ToDtype(scale=True) pair.
Common mistakes
Augmenting the validation set. Your metric then changes between runs and cannot be compared. Fix: separate train_tf and eval_tf, as above.
scale=(0.08, 1.0) on small objects. That is the ImageNet default for RandomResizedCrop, and it can crop to 8% of the area. That is enough to cut your defect out of the frame while keeping its label. Fix: raise the lower bound to 0.5 or 0.7 for small-object tasks.
Rotating without thinking about the corners. A rotated square leaves black wedges, and the model can learn "black wedge means augmented means class such-and-such". Fix: rotate then centre-crop, or use expand=False with a fill value that matches your background.
Doing augmentation on the GPU inside the training loop by habit. For small images CPU workers keep up fine, and moving augmentation to the GPU can leave it idle waiting instead. Fix: measure. See is the GPU waiting for data.
Try it yourself
Add v2.RandomVerticalFlip(p=0.5) to the arrow experiment. Predict the result before running. Then reason about why the answer differs from the horizontal case, given how the arrows are drawn.
What to learn next
- Mixup and CutMix — augmentations that change the label on purpose, in a controlled way.
- RandAugment and AutoAugment — letting a search pick the transformation list instead of you.
- Data augmentation — the same idea across text, audio and tabular data.
Researcher — Mathematics and papers.
Augmentation as an invariance prior
Let $G$ be a group of transformations acting on inputs. Assume the task is invariant: $y(g \cdot x) = y(x)$ for all $g \in G$. Augmented training minimises
$$ \hat{\mathcal{R}}{\text{aug}}(\theta) = \frac{1}{n}\sum{i=1}^{n} \mathbb{E}_{g \sim \mu_G} \big[ \ell(f_\theta(g \cdot x_i),\, y_i) \big] $$
- $G$ — the transformation group (flips, translations, rotations).
- $\mu_G$ — the sampling distribution over $G$, which your
degrees=andscale=arguments define. - $\ell$ — the per-example loss.
- $g \cdot x$ — the transformed image.
Chen, Dobriban and Lee (2020) (NeurIPS / JMLR) formalise this as averaging over group orbits. They prove it produces variance reduction against the unaugmented estimator. Two consequences follow directly. The gain is largest where the model is not already invariant. Augmenting with translations helps a fully-connected network far more than a convolutional one. A CNN is approximately translation-equivariant by construction. And if the invariance assumption is false, the objective above minimises the wrong risk. That is what the arrow experiment measures.
Invariance is not free, and CNNs have less of it than advertised
Zhang (2019), Making Convolutional Networks Shift-Invariant Again (ICML), shows strided convolution and pooling violating the Nyquist criterion. A one-pixel shift can then flip a prediction. Anti-aliased downsampling — blur before stride — restores much of the invariance and improves accuracy at the same time.
Engstrom et al. (2019), Exploring the Landscape of Spatial Robustness (ICML), search over rotations and translations adversarially and find standard classifiers highly vulnerable. Two of their findings matter for practice. Data augmentation gives relatively small robustness against worst-case spatial perturbations. And first-order methods cannot reliably find those worst cases, so spatial robustness behaves unlike $\ell_p$ robustness. Honest evaluation needs a grid search over the transformation.
The practical reading: augmentation buys average-case tolerance, not guarantees.
Interpolation, padding, and the artefacts you are also teaching
Geometric transforms resample, and resampling has a filter. Three details with measurable consequences:
- Antialiasing. Downsampling without a low-pass filter aliases high frequencies. In torchvision v2,
antialias=Trueis the default for tensor inputs on resize-like operations, and matches PIL's behaviour. Mismatched settings between training and evaluation are a real and silent source of accuracy loss. - Padding mode. Zero padding after rotation injects a constant border that correlates with the fact of augmentation. Reflect or replicate padding removes that particular shortcut.
- Chained resampling. Compose a rotation, a crop and a resize as separate operations and you interpolate three times, compounding blur. Frameworks that compose the affine matrix and resample once produce sharper results.
Learned and searched alternatives
Hand-choosing $\mu_G$ is what AutoAugment and RandAugment replace with a search; see RandAugment and AutoAugment. A parallel line builds the invariance into the architecture instead of the data. Cohen and Welling (2016), Group Equivariant Convolutional Networks, give exact rotation equivariance by construction. That is the right choice when the symmetry is genuinely exact, as in microscopy and remote sensing.
Papers
- Chen, Dobriban and Lee, A Group-Theoretic Framework for Data Augmentation, 2020 — arxiv.org/abs/1907.10905
- Zhang, Making Convolutional Networks Shift-Invariant Again, 2019 — arxiv.org/abs/1904.11486
- Engstrom et al., Exploring the Landscape of Spatial Robustness, 2019 — arxiv.org/abs/1712.02779
- Cohen and Welling, Group Equivariant Convolutional Networks, 2016 — arxiv.org/abs/1602.07576
- Shorten and Khoshgoftaar, A survey on Image Data Augmentation for Deep Learning, Journal of Big Data 2019
What to learn next
- Mixup and CutMix — augmentations that change the label on purpose, in a controlled way.
- RandAugment and AutoAugment — letting a search pick the transformation list instead of you.
- Data augmentation — the same idea across text, audio and tabular data.