Training Vision Models

Test-time augmentation

Test-time augmentation runs the same model on several altered copies of one test image and averages the answers, buying a small accuracy gain and better-behaved confidence for several times the inference cost.

On this page 5
  1. Why it works at all
  2. The rule you already know
  3. What you should expect to get
  4. Remember this
  5. 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.

Test-time augmentation means showing the model several versions of one photo and averaging its answers.

You are trying to read a faded shop sign across the road. You read it once. Then you step left and read it again. Then you move closer and squint. Four looks, four readings, and you go with what you saw most.

You did not become a better reader between looks. You collected more evidence about the same sign.

That is exactly what test-time augmentation does, and the word for it in short is TTA.

Why it works at all

A trained model is not perfectly steady. Nudge a photo two pixels to the left and the answer can wobble. Mirror it and the confidence changes.

That wobble is a weakness during normal use. TTA turns it into a resource. Each altered version makes a slightly different mistake, and mistakes that disagree with each other partly cancel when you average.

   one test photo
        |
   +----+----+-------+
   |    |    |       |
 as-is  L-R  U-D   both ways
   |    |    |       |
   v    v    v       v
    [ the same model, four times ]
   |    |    |       |
   +----+--+-+-------+
             |
      average the four answers
             |
        final answer

Nothing is trained. It is the same model, four times over, on four versions of one picture.

The rule you already know

The alterations must be ones that do not change the true answer. Exactly the rule from geometric augmentation, and it bites harder here.

Mirror a photo of a cat: still a cat, safe. Mirror a photo of an arrow pointing left and it points right. You have added a confidently wrong vote. During training a bad augmentation weakens learning. At test time it directly corrupts the prediction you are about to act on.

What you should expect to get

Be sceptical of anyone promising a lot. In the developer section, four flipped views on a hard task moved accuracy from 0.7400 to 0.7425. A quarter of a point, for four times the compute.

The measurable improvement was in overconfidence. The single-view model claimed 87 percent certainty while being right 74 percent of the time. Averaging brought the claimed certainty down towards the truth without losing accuracy.

So TTA is worth considering when being right matters more than being fast. It also helps when you use the stated confidence. It is a poor choice for a real-time camera. Four times the cost per frame is not on offer.

Remember this

  • TTA runs the same model on several altered copies of one image and averages the answers.
  • The alterations must not change the true label, or you are averaging in lies.
  • Expect a small accuracy gain and better-behaved confidence, at several times the inference cost.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch numpy

Written and run against torch 2.13.0 (CPU) and numpy 2.2.6. No downloads. Trains a small CNN from scratch on noisy synthetic shapes; the whole script runs in a few seconds.

The shapes are square, disc, ring and cross. All four are symmetric left-right and up-down, so flips genuinely preserve the label. That is deliberate — the experiment is about aggregation, not about a broken transform.

Measuring TTA honestly, view by view

tta.py
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.data import TensorDataset, DataLoader

rng = np.random.default_rng(0); torch.manual_seed(0)

def draw(cls):                      # four shapes a flip cannot turn into each other
    img = rng.normal(0.4, 0.30, (32, 32))          # heavy noise: the task is genuinely hard
    r, c = rng.integers(9, 23, 2); s = int(rng.integers(4, 7))
    yy, xx = np.ogrid[:32, :32]; d = np.hypot(yy - r, xx - c)
    m = [(abs(yy-r) < s) & (abs(xx-c) < s), d < s, abs(d-s) < 1.2,
         (abs(yy-r) < 1.2) | (abs(xx-c) < 1.2)][cls]
    img[m] += 0.55
    return np.clip(img, 0, 1).astype("float32")

def make(n):
    lab = rng.integers(0, 4, n)
    return torch.tensor(np.stack([draw(c) for c in lab])).unsqueeze(1), torch.tensor(lab)

Xtr, ytr = make(2000); Xte, yte = make(800)
net = nn.Sequential(nn.Conv2d(1,16,3,padding=1), nn.BatchNorm2d(16), nn.ReLU(), nn.MaxPool2d(2),
                    nn.Conv2d(16,32,3,padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2),
                    nn.Flatten(), nn.Linear(32*8*8, 4))
opt = torch.optim.AdamW(net.parameters(), lr=1e-3)
dl = DataLoader(TensorDataset(Xtr, ytr), batch_size=64, shuffle=True)
for _ in range(12):
    net.train()
    for xb, yb in dl:
        opt.zero_grad(); F.cross_entropy(net(xb), yb).backward(); opt.step()
net.eval()

def shift(t, dy, dx): return torch.roll(t, shifts=(dy, dx), dims=(-2, -1))

VIEWS = [("original", lambda t: t), ("h-flip", lambda t: t.flip(-1)), ("v-flip", lambda t: t.flip(-2)),
         ("both flips", lambda t: t.flip(-1).flip(-2)),
         ("shift +2,+2", lambda t: shift(t, 2, 2)),  ("shift -2,-2", lambda t: shift(t, -2, -2)),
         ("shift +2,-2", lambda t: shift(t, 2, -2)), ("shift -2,+2", lambda t: shift(t, -2, 2))]

with torch.no_grad():
    ps = [(n, F.softmax(net(f(Xte)), 1)) for n, f in VIEWS]
for n, p in ps:
    print(f"{n:<12} acc {(p.argmax(1)==yte).float().mean():.4f}  mean confidence {p.max(1).values.mean():.3f}")
for k in (4, 8):
    avg = torch.stack([p for _, p in ps[:k]]).mean(0)
    print(f"TTA over {k}   acc {(avg.argmax(1)==yte).float().mean():.4f}  mean confidence {avg.max(1).values.mean():.3f}")
Output
original     acc 0.7400  mean confidence 0.873
h-flip       acc 0.7337  mean confidence 0.866
v-flip       acc 0.7412  mean confidence 0.876
both flips   acc 0.7312  mean confidence 0.875
shift +2,+2  acc 0.7175  mean confidence 0.837
shift -2,-2  acc 0.7188  mean confidence 0.858
shift +2,-2  acc 0.7113  mean confidence 0.856
shift -2,+2  acc 0.7025  mean confidence 0.847
TTA over 4   acc 0.7425  mean confidence 0.861
TTA over 8   acc 0.7412  mean confidence 0.845

Deterministic with the seeds set.

What that table is really saying

The four flip views disagree, on a task where they should not. 0.7400, 0.7337, 0.7412, 0.7312. Every shape is symmetric, so a perfectly flip-invariant model would give four identical rows. The spread is the model's instability, and it is the raw material TTA works with.

Averaging four views gained 0.0025. From 0.7400 to 0.7425 — two extra correct answers out of 800, for four times the inference cost. Report that honestly and let the product decide.

Adding four more views made it slightly worse. 0.7425 down to 0.7412. The shift views individually score 0.70 to 0.72, below the original's 0.74. Averaging in weaker voters dragged the result down. More views is not better. Views that are individually good and mutually different are better.

Confidence fell from 0.873 to 0.845 while accuracy stayed put. The single-view model claimed 87.3 percent certainty at 74.0 percent accuracy — a thirteen-point overconfidence gap. Eight-view TTA narrowed it by about three points at no accuracy cost. If you threshold on confidence, that is the effect you are buying.

Line by line

torch.roll shifts pixels and wraps them around the edge. Real TTA would crop or pad instead. Roll keeps the script dependency-free, and its wrap-around is why the shift views score lower.

F.softmax(...) before averaging matters. Averaging probabilities and averaging raw scores are different operations and give different answers. Probabilities are the safer default because raw scores are not on a comparable scale across views.

net.eval() once, before all views. In training mode, BatchNorm would recompute statistics per view and every row would be measuring something else.

torch.stack([...]).mean(0) is the plain average. Weighted aggregation learns those weights on a validation set. The researcher section covers it, and it does better than this.

When TTA is worth it

SituationVerdict
Offline batch scoring, accuracy matters mostWorth trying
Medical or safety review with a human in the loopOften worth it, for the confidence behaviour
A leaderboard or competitionStandard practice
Real-time video, per-frame budgetRarely; the cost multiplies directly
Any task where flips change the labelDo not, unless you pick label-safe views

Common mistakes

Using the training augmentation list at test time. Training uses harsh transforms on purpose. At test time each view must be a plausible version of the input. Fix: a short, gentle, separate view list.

Including views that break the label. A horizontal flip on digits, text, or left-versus-right classes adds confident wrong votes. Fix: run each view alone first, as the table above does, and drop any that scores far below the original.

Averaging predicted classes instead of probabilities. Majority voting over four views throws away all the confidence information and ties often. Fix: average probabilities.

Reporting TTA accuracy against a non-TTA baseline tuned differently. Fix: same weights, same test set, only the view list changing.

Forgetting the cost in your latency budget. Eight views is eight times the compute per image. Fix: measure it before promising it. See latency and throughput.

Try it yourself

Swap the four shape classes for the two arrow classes from geometric augmentations. Keep the horizontal-flip view, and watch TTA do worse than a single forward pass. That is the same error as the training-time one, except here it damages predictions you are about to act on.

What to learn next

Researcher — Mathematics and papers.

The estimator

For a model $f$, an input $x$, and a set of label-preserving transforms ${t_1, \dots, t_K}$, plain TTA computes

$$ \hat{p}(x) = \frac{1}{K}\sum_{k=1}^{K} f\big(t_k(x)\big) $$

  • $t_k$ — the $k$-th view transform, with $t_1$ usually the identity.
  • $f(\cdot)$ — the model's output probability vector.
  • $K$ — the number of views, and the multiplier on inference cost.

This is an ensemble over inputs rather than over models. It shares the variance-reduction argument of ensembling and costs one model instead of $K$. But the members are highly correlated, being the same weights. So the achievable gain is much smaller than a true deep ensemble's.

Ashukha et al. (2020) (ICLR) make this quantitative with the deep ensemble equivalent score. It reports how many independently trained networks a technique matches in test log-likelihood. Many sophisticated methods score equivalent to only a handful. It is the right yardstick for asking whether your $K$-fold inference bill bought anything.

Averaging is not the best aggregator

Shanmugam et al. (2021), Better Aggregation in Test-Time Augmentation (ICCV), attack the mean directly. Their central empirical observation: even when TTA produces a net accuracy improvement, it changes many correct predictions into incorrect ones. The net gain hides two flows in opposite directions, and the plain average is what allows the harmful flow.

They analyse when the simple average is suboptimal, then learn a per-class aggregation function on validation data. They report consistent improvements over the average across models, datasets and augmentations.

Two consequences for practice:

  • Report the flip matrix, not only the net delta. Count correct-to-incorrect and incorrect-to-correct separately. Take a +0.2% net gain that flips 3% of predictions in each direction. Its risk profile differs from one that flips 0.3%.
  • Learn the weights if you have a validation set. Uniform weighting is a default, not an optimum, and the experiment above shows it being dragged down by weak views.

History, and why the old numbers were bigger

TTA is as old as modern ImageNet practice. Krizhevsky et al. (2012) averaged predictions over ten crops — four corners, centre, and their mirrors. Simonyan and Zisserman (2015) used dense evaluation and multi-crop for VGG, reporting meaningful gains.

The gains were larger then for a reason that no longer holds. Those models trained with far weaker augmentation, so they were much less invariant to crops and flips. That left more instability for TTA to average away. A model trained today with RandAugment and CutMix already has those invariances. The residual left for TTA to recover is small. Expect smaller TTA gains on better-trained models, and treat old reported numbers as a different regime.

Relationship to test-time adaptation

TTA and test-time adaptation are different, and the abbreviations collide. Augmentation changes the input and leaves the weights alone. Adaptation changes the model using unlabelled test data — see domain adaptation for vision for AdaBN and Tent. They compose: adapt the normalisation statistics to the new domain, then average over views. Combining them multiplies neither cost nor risk in an obvious way, and few papers evaluate the combination carefully.

Papers

  • Krizhevsky, Sutskever and Hinton, ImageNet Classification with Deep Convolutional Neural Networks, NeurIPS 2012
  • Simonyan and Zisserman, Very Deep Convolutional Networks for Large-Scale Image Recognition, 2015 — arxiv.org/abs/1409.1556
  • Shanmugam et al., Better Aggregation in Test-Time Augmentation, 2021 — arxiv.org/abs/2011.11156
  • Ashukha et al., Pitfalls of In-Domain Uncertainty Estimation and Ensembling in Deep Learning, 2020 — arxiv.org/abs/2002.06470

What to learn next