Training Vision Models

RandAugment and AutoAugment

AutoAugment searched for a good list of image transformations and wrote it down, RandAugment threw the search away and left two knobs, and on a small clean dataset both can cost you accuracy.

On this page 6
  1. Why the chef came first
  2. What RandAugment changed
  3. The strength knob matters more than the list
  4. Where you have seen the effect
  5. Remember this
  6. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

RandAugment picks a couple of image edits at random and applies them at a strength you choose.

Think about cooking with a spice rack. One approach is to hire a chef. She tastes four thousand combinations over six months. She hands you an exact recipe: a pinch of this, then that.

The other approach is to grab any two jars from the rack and use a fixed amount of each. It sounds careless. It turns out to taste nearly as good, and it takes a second to explain.

AutoAugment is the chef. RandAugment is grabbing two jars.

Why the chef came first

Once people knew augmentation helped, an obvious question followed. Which transformations, in what order, at what strength, how often?

Nobody knew. The answer differs by dataset, and testing a combination means training a model. AutoAugment turned that into a search problem: try policies, train small models, keep what works. The result was a fixed recipe you copy into your code.

Two problems appeared. The search cost thousands of GPU hours. And the recipe was tuned on small models and small datasets. Applied to large ones, it was no longer right.

What RandAugment changed

RandAugment threw the search away and left two knobs.

   knob one:  how many edits to apply to each photo   (usually two)
   knob two:  how strong each edit is                 (a number; higher is harsher)

   for each photo:
       pick that many edits at random from a fixed list
       apply every one of them at that strength

The fixed list is ordinary stuff: rotate, shear, translate, change brightness, change contrast, sharpen, posterise, invert, equalise.

Because there are only two numbers, an ordinary grid search finds them. On your own dataset, inside your own training run. That is the real contribution. It replaced a research project with a hyperparameter.

The strength knob matters more than the list

Here is the finding that most tutorials skip. The right strength depends on how big your model is and how much data you have.

Big model, millions of images, long training: crank it up. The model has capacity to spare and time to learn through the distortion.

Small dataset, short fine-tune, clean well-framed photos: turn it down, or off. Harsh distortion on a small clean dataset destroys signal you cannot afford to lose.

The developer section measures this on a real face dataset. Default-strength RandAugment makes the model worse. Not by much, and consistently.

Where you have seen the effect

Every headline image model of the last several years used something from this family during training. It is one of the quiet reasons published accuracy numbers kept rising while the architectures stopped changing much.

Remember this

  • AutoAugment is a searched, fixed recipe. RandAugment is two knobs: how many edits, how strong.
  • The right strength grows with model size and dataset size.
  • On a small clean dataset these can lose you accuracy. Measure before you keep them.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch torchvision scikit-learn

Written and run against torch 2.13.0 (CPU), torchvision 0.28.0, scikit-learn 1.7.2.

All four members of the family live in torchvision.transforms.v2. They expect image tensors — uint8 is the natural input, and float tensors are accepted too.

What is available, and that it does something

which_ones.py
import torch
from torchvision.transforms import v2, AutoAugmentPolicy

torch.manual_seed(0)
img = (torch.rand(3, 32, 32) * 255).to(torch.uint8)      # uint8 is what these expect

for name, t in [("RandAugment(2, 9)", v2.RandAugment(num_ops=2, magnitude=9)),
                ("AutoAugment IMAGENET", v2.AutoAugment(AutoAugmentPolicy.IMAGENET)),
                ("AutoAugment CIFAR10", v2.AutoAugment(AutoAugmentPolicy.CIFAR10)),
                ("TrivialAugmentWide", v2.TrivialAugmentWide()),
                ("AugMix", v2.AugMix())]:
    out = t(img)
    print(f"{name:<22} -> dtype {out.dtype}, shape {tuple(out.shape)}, pixels changed: {not torch.equal(out, img)}")
Output
RandAugment(2, 9)      -> dtype torch.uint8, shape (3, 32, 32), pixels changed: True
AutoAugment IMAGENET   -> dtype torch.uint8, shape (3, 32, 32), pixels changed: True
AutoAugment CIFAR10    -> dtype torch.uint8, shape (3, 32, 32), pixels changed: True
TrivialAugmentWide     -> dtype torch.uint8, shape (3, 32, 32), pixels changed: True
AugMix                 -> dtype torch.uint8, shape (3, 32, 32), pixels changed: True

Note AutoAugmentPolicy has three members: IMAGENET, CIFAR10 and SVHN. They are three different searched recipes, and picking the wrong one is a real mistake. The CIFAR policy was tuned on 32-pixel images.

What the magnitude knob actually does

magnitude.py
import torch
from sklearn.datasets import fetch_olivetti_faces
from torchvision.transforms import v2

faces = fetch_olivetti_faces()
img = torch.tensor(faces.images[0]).mul(255).to(torch.uint8).unsqueeze(0).repeat(3, 1, 1)  # one real face

print(f"{'magnitude':>10} {'mean |pixel change| (0-255)':>29}")
for mag in (1, 5, 9, 17, 25, 30):
    torch.manual_seed(0)                       # same random draws at every magnitude
    aug = v2.RandAugment(num_ops=2, magnitude=mag)
    d = [(aug(img).float() - img.float()).abs().mean().item() for _ in range(30)]
    print(f"{mag:>10} {sum(d)/len(d):>29.1f}")

try:
    v2.RandAugment(magnitude=31)(img)
except Exception as e:
    print("\nmagnitude=31 ->", type(e).__name__, ":", str(e)[:60])
Output
 magnitude   mean |pixel change| (0-255)
         1                           7.8
         5                          22.8
         9                          40.3
        17                          56.8
        25                          68.4
        30                          74.5

magnitude=31 -> IndexError : index 31 is out of bounds for dimension 0 with size 31

At magnitude 30, the average pixel moves by 74 levels out of 255. Roughly a third of the dynamic range, on every pixel, on every training image. That is not a gentle nudge.

The IndexError is worth knowing. magnitude indexes into a table of num_magnitude_bins values, default 31, so the legal range is 0 to 30. The error message does not mention magnitude at all, and people lose an hour to it.

Does it help? Test it, do not assume

Fine-tuning ResNet-18 on 40-way face identification, six photos per person, with and without RandAugment.

does_it_help.py
import torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import fetch_olivetti_faces
from torch.utils.data import TensorDataset, DataLoader
from torchvision.transforms import v2
from torchvision.models import resnet18, ResNet18_Weights

faces = fetch_olivetti_faces()
raw = torch.tensor(faces.images).mul(255).to(torch.uint8).unsqueeze(1).repeat(1, 3, 1, 1)
raw = v2.functional.resize(raw, [112, 112])       # stay uint8: RandAugment wants integer pixels
y = torch.tensor(faces.target)
tr = torch.cat([torch.arange(i*10, i*10+6) for i in range(40)])
te = torch.cat([torch.arange(i*10+6, i*10+10) for i in range(40)])

to_float = v2.Compose([v2.ToDtype(torch.float32, scale=True),
                       v2.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])

def run(aug, epochs=6):
    torch.manual_seed(0)
    m = resnet18(weights=ResNet18_Weights.DEFAULT); m.fc = nn.Linear(512, 40)
    opt = torch.optim.AdamW(m.parameters(), lr=1e-4)
    dl = DataLoader(TensorDataset(raw[tr], y[tr]), batch_size=16, shuffle=True)
    for _ in range(epochs):
        m.train()
        for xb, yb in dl:
            if aug is not None: xb = aug(xb)      # augment while still uint8
            opt.zero_grad()
            F.cross_entropy(m(to_float(xb)), yb).backward()
            opt.step()
    m.eval()
    with torch.no_grad():
        p = torch.cat([m(to_float(raw[te][i:i+40])) for i in range(0, 160, 40)]).argmax(1)
    return (p == y[te]).float().mean().item()

print("no augmentation      :", round(run(None), 3))
print("RandAugment mag 9    :", round(run(v2.RandAugment(num_ops=2, magnitude=9)), 3))
print("RandAugment mag 25   :", round(run(v2.RandAugment(num_ops=2, magnitude=25)), 3))
Output
no augmentation      : 0.969
RandAugment mag 9    : 0.944
RandAugment mag 25   : 0.894

Deterministic with the seed set. On a laptop CPU this took about 50 seconds in total, most of it the three fine-tuning runs.

Sit with that result

RandAugment made it worse, and harder settings made it worse still. 0.969, then 0.944, then 0.894. The ordering is monotonic in magnitude. That is what you expect when augmentation destroys signal rather than adding variety.

This does not contradict the paper. Cubuk et al. report gains at ImageNet scale: 1.28 million images, hundreds of epochs, large models. They designed RandAugment so strength could be tailored to model and dataset size. Policies tuned on small proxies did not transfer to large ones.

These faces are the opposite regime. 240 training images, six epochs, a small model, and photographs that are already cropped, centred and consistently lit. There is little overfitting to regularise away and plenty of fine detail to destroy.

The correct conclusion is not "RandAugment is bad". Augmentation strength is a hyperparameter with a real optimum. On your data that optimum may be zero. Nobody can tell you which without running the three lines above.

Common mistakes

Applying it after normalisation. These transforms include operations like Posterize and Equalize that assume pixel-like values. Run them on uint8 images before ToDtype and Normalize, as the script above does.

Using AutoAugmentPolicy.CIFAR10 on 224-pixel photos. The searched magnitudes were tuned for 32-pixel images. Fix: use IMAGENET for photographs, and prefer RandAugment when your data resembles neither.

Augmenting the evaluation set. Your accuracy then wobbles between runs for no reason. Fix: separate train and eval transform pipelines.

Assuming the search transfers. An AutoAugment policy found for CIFAR is not a general truth about images. Fix: treat any policy as a starting point and tune magnitude on your own validation set.

Stacking RandAugment on top of a full geometric pipeline. RandAugment already contains rotation, shear and translation. Adding your own on top can compound to distortions neither of you intended. Fix: pick one source of geometry.

Try it yourself

Add v2.TrivialAugmentWide() as a fourth row in the last experiment. It has no knobs at all — it samples one operation and one strength uniformly per image. Compare it against your best hand-tuned magnitude, and note how much tuning effort it saved or cost.

What to learn next

Researcher — Mathematics and papers.

AutoAugment: the search formulation

Cubuk et al. (2019), AutoAugment: Learning Augmentation Strategies from Data (CVPR), define a policy as 5 sub-policies. Each sub-policy is a pair of operations, and each operation carries a probability and a magnitude. With 16 operations and discretised probability and magnitude, the space is roughly $10^{32}$ policies. An RNN controller trained by reinforcement learning searched it, using validation accuracy of a child model as reward.

The cost was the problem. Searching required training thousands of child models — thousands of GPU-hours per dataset. Lim et al. (2019), Fast AutoAugment, cut this by density matching rather than child-model training; Ho et al. (2019), Population Based Augmentation, learned a schedule rather than a fixed policy.

Cubuk et al. (2020) (CVPR Workshops; also NeurIPS 2020) reduce the space to two integers:

$$ |\mathcal{S}| = K^{N} $$

  • $K$ — the number of available operations (14 in the paper).
  • $N$ — operations applied per image.
  • $M$ — a single magnitude shared by all operations, held constant.

With $N \le 3$ this is small enough for a naive grid search inside the normal training run. The separate proxy-search phase disappears.

Two findings from that paper deserve emphasis:

The magnitude schedule did not matter. Random, constant, linearly increasing and random-with-increasing-bound magnitudes all worked comparably. They chose constant, for having one hyperparameter.

Optimal strength grows with model and dataset size. This is the paper's core critique of proxy-based search. Policies optimised on small models and reduced datasets are systematically too weak for the large-scale setting. The experiment in the developer section is the same effect running the other way. An ImageNet-calibrated default magnitude is too strong for a small clean dataset.

RandAugment matched AutoAugment and Fast AutoAugment on ResNet-50, and exceeded them on larger models. Their strongest setting reached 85.0% ImageNet accuracy.

The uncomfortable follow-up

Müller and Hutter (2021), TrivialAugment (ICCV oral), removed the remaining knobs. TrivialAugment applies exactly one operation per image, with the operation and its strength both drawn uniformly at random. It is parameter-free, and they report it matching or beating the tuned methods.

That is a strong negative result for the whole search programme. One uniformly sampled operation is competitive with a policy costing thousands of GPU-hours. Most of what the search found was therefore not transferable structure.

AugMix, for a different objective

Hendrycks et al. (2020) (ICLR) target corruption robustness rather than clean accuracy. AugMix samples several augmentation chains and mixes their outputs convexly. It then adds a Jensen-Shannon consistency loss, so the clean image and its augmented versions embed similarly:

$$ \mathcal{L}{\text{JS}} = \frac{1}{3}\Big( \mathrm{KL}(p{\text{orig}} \Vert \bar{p}) + \mathrm{KL}(p_{\text{aug1}} \Vert \bar{p}) + \mathrm{KL}(p_{\text{aug2}} \Vert \bar{p}) \Big) $$

  • $p_{\text{orig}}, p_{\text{aug1}}, p_{\text{aug2}}$ — output distributions for the clean image and two augmented versions.
  • $\bar{p}$ — their mean.
  • $\mathrm{KL}$ — Kullback-Leibler divergence.

They report substantially improved robustness and uncertainty on corruption benchmarks. v2.AugMix in torchvision implements the image operation only. The consistency loss is yours to add. Without it you have the weaker half of the method. This is a common and quiet omission.

Papers

What to learn next