Skip to content

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_aug
test_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 = 5
score_name = valid/bal_acc
tester_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 -> 1088 over 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:

examples/classification/dataloaders/data_utils/balanced_sampler.py
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,
    )