Optimisers, Schedulers and the Training Loop
TorchMetrics
TorchMetrics wraps each metric in an object that accumulates correctly across batches and devices, replacing the hand-rolled counters that so often lie.
- 6 min read
- 3 reading levels
- Published
Read these first
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
TorchMetrics is a library of ready-made scorekeepers: you feed each one predictions batch by batch, and it keeps the running totals needed to report the true final score.
Think of a cricket scoreboard operator. After every ball, you tell him what happened. He maintains the totals — runs, wickets, overs — and at any moment can state the score exactly.
You never hand him "the average of the last few balls". You hand him events; the totalling is his job, and he does it right.
Why it exists
The previous lesson showed how batch-by-batch averaging goes wrong for loss. For fancier scores it gets worse. Take F1, a score balancing how many positives you found against how many of your alarms were real. That one cannot be averaged across batches at all. Each batch's F1 is a ratio, and averaging ratios of different batches produces a number that corresponds to nothing.
The safe recipe is always: keep raw tallies, compute the score once at the end. TorchMetrics packages that recipe for over a hundred metrics, so you stop hand-rolling it.
How it works
batch 1 predictions ──→ ┌────────────────┐
batch 2 predictions ──→ │ metric object │ ──→ .compute() → one true score
batch 3 predictions ──→ │ (keeps tallies) │
└────────────────┘Three verbs run everything: update (feed a batch), compute (report the score), reset (clear the tallies for the next epoch).
Where you have seen this
Scoreboards, odometers, electricity meters — every trustworthy running total works this way: record raw events, derive the figure from totals. TorchMetrics is that discipline applied to model scores.
Remember this
- Metric objects accumulate tallies; you never average batches yourself.
- The lifecycle is update → compute → reset, once per epoch.
- Scores like F1 cannot be batch-averaged — this library exists because of them.
What to learn next
- Early stopping and keeping the best model — acting on the honest scores you now have.
- Model evaluation — choosing which metric deserves to be the headline.
- Writing a validation loop that reports the truth — the hand-rolled foundation this library packages.
Developer — Code and libraries.
Setup
pip install torch torchmetricsTorchMetrics is a small, pure-Python install (about 1 MB plus a tiny helper package). Outputs captured with torchmetrics 1.9.0 and torch 2.5.1 on CPU.
The three-verb lifecycle
import torch
from torchmetrics.classification import MulticlassAccuracy, MulticlassF1Score
torch.manual_seed(0)
accuracy = MulticlassAccuracy(num_classes=3, average="micro")
f1 = MulticlassF1Score(num_classes=3, average="macro")
for _ in range(4): # pretend: four validation batches
logits = torch.randn(25, 3) # metrics accept raw logits directly
target = torch.randint(0, 3, (25,))
accuracy.update(logits, target)
f1.update(logits, target)
print(f"accuracy over all 100 samples: {accuracy.compute():.4f}")
print(f"macro F1 over all 100 samples: {f1.compute():.4f}")
accuracy.reset() # forget everything before the next epoch
print(f"after reset, updates seen: {accuracy.update_count}")accuracy over all 100 samples: 0.3900 macro F1 over all 100 samples: 0.3843 after reset, updates seen: 0
Random predictions against three classes land near 33% — the 0.39 here is ordinary luck on 100 samples, a useful reminder of how noisy small validation sets are.
Walkthrough
MulticlassAccuracy(num_classes=3, ...) — since TorchMetrics v1, classification metrics are explicit about their task: Binary*, Multiclass*, Multilabel* classes (or a task= argument on wrapper classes). The num_classes is required; it sizes the internal tally tensors.
update(logits, target) — accepts raw logits, probabilities, or predicted class indices, and detects which you passed from the shape and values. Handing logits straight from the model, before any softmax or argmax, is the intended use.
average="micro" vs "macro" — micro pools every sample into one tally, so frequent classes dominate. Macro computes the score per class and averages the classes equally, so rare classes count as much as common ones. On imbalanced data the two can disagree wildly; choosing between them is a decision about what you care about, not a technicality.
compute() — derives the score from the tallies. Calling it mid-epoch is allowed (it computes over everything so far), which is handy for progress displays.
Metrics are nn.Modules. They live on a device and must match their inputs: metric.to(device) for GPU validation. They also slot into a ModuleDict inside your model class, and checkpoint with it.
Common mistakes
Forgetting reset() between epochs. The tallies happily span epochs, blending stale results into every later score. Epoch 2's "improvement" may be arithmetic fog. Reset after logging, every epoch.
Averaging compute() calls per batch. Calling compute() on each batch and averaging the results reintroduces exactly the bug the library removes. One update per batch, one compute per epoch.
Device mismatch. A metric left on CPU while GPU tensors flow in raises a device error — the same family as in device mismatch errors. Move the metric where the tensors are.
Recreating the metric inside the loop. A fresh object each batch has no memory; its final answer covers one batch. Construct metrics once, next to the model.
Try it yourself
Set both metrics to average="macro", then make the data imbalanced — for example, target = torch.where(torch.rand(25) < 0.9, 0, torch.randint(1, 3, (25,))). Watch accuracy and macro F1 pull apart, and explain which one an imbalanced-fraud-detection team should report. More on that trade-off in imbalanced data.
What to learn next
- Early stopping and keeping the best model — acting on the honest scores you now have.
- Model evaluation — choosing which metric deserves to be the headline.
- Writing a validation loop that reports the truth — the hand-rolled foundation this library packages.
Researcher — Mathematics and papers.
State design
Each metric declares its state via add_state(name, default, dist_reduce_fx) — tensors updated by update() and reduced across processes at compute() time. For MulticlassAccuracy the states are class-wise tp/fp/tn/fn count tensors with dist_reduce_fx="sum"; for AUROC the state is the retained score/target lists (or a binned approximation), because AUC is a pairwise statistic without fixed-size sufficient statistics.
This is the formal decomposability split: metrics whose dataset value is $f(\sum_b s_b)$ for per-batch statistics $s_b$ (accuracy, confusion-based scores, mean losses) accumulate O(1) state; metrics defined over sample pairs or global ranks (AUROC, Spearman, calibration error with adaptive bins) need O(n) state or approximation.
Averaging conventions
For per-class scores $M_c$ with supports $n_c$:
$$\text{macro} = \frac{1}{C}\sum_c M_c, \qquad \text{weighted} = \sum_c \frac{n_c}{n} M_c, \qquad \text{micro: pool all samples, then score}$$
- $C$ — number of classes; $n_c$ — samples of class $c$.
For multiclass accuracy, micro-averaging coincides with plain accuracy. For F1, micro-F1 equals accuracy in single-label settings — a frequent source of "why are these identical" confusion.
Distributed correctness
Under DDP, compute() triggers the declared reductions via the process group, so per-rank shards yield the exact global metric — including the uneven-shard case that breaks naive per-rank averaging (see the distributed caveat in the validation lesson). sync_on_compute=True is the default; metric(preds, target) (forward) additionally returns the batch-local value while updating global state.
Reference
- Detlefsen et al. (2022), TorchMetrics — Measuring Reproducibility in PyTorch, Journal of Open Source Software 7(70) — the design paper.
- Documented at lightning.ai/docs/torchmetrics; version-pin in projects (
torchmetrics==1.9.*), as the v0→v1 transition renamed most classification classes and older tutorials still show the pre-2023 API.
What to learn next
- Early stopping and keeping the best model — acting on the honest scores you now have.
- Model evaluation — choosing which metric deserves to be the headline.
- Writing a validation loop that reports the truth — the hand-rolled foundation this library packages.