Classification (the realistic one)¶
Synthetic 32x32 noisy shapes, five classes, class-imbalanced. CPU, no downloads, ~17 s.
gatle-ignite train --config=examples/classification/configs/classification_v0.py
python examples/classification/scripts/checks.py # 13 checks, each falsifiable
Synthetic is the smallest task that works. This is the one to crib from for a real project: it wires up the parts an actual classification run needs: augmentation, a class-balancing sampler, top-k and per-class metrics, warmup, and all three engines.
Nothing beyond the walk: model · prep_batch · loss · optimizer · checkpoints · logging
| step | what this example changes |
|---|---|
| dataset | aug_name = examples.classification.augmentation.shapes_augtest_ds_name = examples.classification.dataloaders.shapes_dataset |
| metrics | val_metrics: acc (accuracy), bal_acc (per_class), top2 (topk)tester_metrics: acc (accuracy), bal_acc (per_class), top2 (topk)every_test = 5score_name = valid/bal_acctester_score_name = test/bal_acc |
| running | grad_clip_norm = 1.0 |
The point: accuracy is the wrong score¶
The train set is imbalanced 4:1 and the test set carries a deployment skew of 400:200:100:40:20. Running with and without the sampler:
| test/acc | test/bal_acc | hbar | vbar | |
|---|---|---|---|---|
| sampler on | 0.542 | 0.688 | 0.95 | 0.85 |
| sampler off | 0.611 | 0.365 | 0.00 | 0.00 |
The worse model has the higher plain accuracy. It bought those points by abandoning two of five
classes entirely: hbar and vbar score exactly 0.00. That is why this example scores on
valid/bal_acc, and it is the whole reason the example exists.
Reproduce it yourself:
SHAPES_NO_SAMPLER=1 gatle-ignite train --config=examples/classification/configs/classification_v0.py
Each wired-up piece can die silently¶
Augmentation that never fires, a sampler that never rebalances, a test engine that never runs: none
of them raise, and accuracy still looks plausible. So checks.py proves each one separately:
- The sampler rebalances: 20 000 draws give
{square .201, circle .197, ring .200, hbar .201, vbar .200}against a raw distribution of[.527, .266, .130, .052, .026]. - Augmentation hits train and not eval: the transform's call count goes
0 -> 1088over one train epoch, then stays at 1088 across a full valid+test pass. - The test engine ran: it fires at epochs 5 and 10 only, and writes a
test_best_result_*checkpoint.
The checks are mutation-tested. Removing the split == "train" guard turns the augmentation checks
red; disabling the sampler turns the balance checks red. A check nobody has watched fail is
indistinguishable from decoration.
The sampler itself is an ordinary dotted component: sampler_params.cls_name in
train_ds_params, exposing get_sampler:
import torch
from torch.utils.data import WeightedRandomSampler
def get_sampler(dataset, num_samples=None, num_classes=None, replacement=True):
"""Sample each class with equal probability, by weighting 1/count.
Reads `dataset.labels`: iterating the dataset would run the augmentation per sample.
"""
labels = torch.as_tensor(getattr(dataset, "labels", None))
if labels.ndim == 0:
raise ValueError("dataset exposes no `labels`; the sampler cannot weight classes")
counts = torch.bincount(labels, minlength=num_classes or int(labels.max()) + 1).clamp(min=1)
weights = (1.0 / counts.float())[labels]
return WeightedRandomSampler(
weights=weights,
num_samples=num_samples or len(labels),
replacement=replacement,
)