Skip to content

MNIST

A realistic task, showing the pieces the synthetic example leaves out: an augmentation module, a label-smoothed loss, a separate test split with its own engine, and gradient clipping.

pip install "gatle-ignite[examples]"   # needs torchvision
# or from a checkout:  pip install -e ".[examples]"
gatle-ignite train --config=examples/mnist/configs/mnist_v0.py

Downloads MNIST to ./data on first run. Reaches ~97% in one epoch on CPU.

Nothing beyond the walk: model · prep_batch · loss · optimizer

step what this example changes
dataset aug_name = examples.mnist.augmentation.mnist_aug
test_ds_name = examples.mnist.dataloaders.mnist_dataset
metrics tester_metrics: acc (accuracy)
every_test = 5
checkpoints resume = True
logging logger_name = ['text', 'pbar']
running grad_clip_norm = 1.0

Augmentation, and who gets it

examples/mnist/augmentation/mnist_aug.py
from torchvision import transforms

from examples.mnist.dataloaders.mnist_dataset import MEAN, STD


class Transformation:
    """A torchvision Compose in a class: the contract asks only for a callable."""

    def __init__(self, degrees=10, translate=0.1):
        self.transform = transforms.Compose(
            [
                transforms.RandomAffine(degrees=degrees, translate=(translate, translate)),
                transforms.ToTensor(),
                transforms.Normalize(MEAN, STD),
            ]
        )

    def __call__(self, image):
        return self.transform(image)

cfg.aug_name builds one transform and the framework passes it to every split. That is a trap: a dataset that applies it blindly will augment the validation set and quietly depress every metric you use to make decisions. So the dataset decides:

examples/mnist/dataloaders/mnist_dataset.py
from torchvision import datasets, transforms

from gatle_ignite import build_dataloader

MEAN, STD = (0.1307,), (0.3081,)


def get_ds(ds_params, transform=None):
    """MNIST. Needs the [examples] extra (torchvision) and downloads on first run."""
    train = ds_params.get("split", "train") == "train"

    plain = transforms.Compose([transforms.ToTensor(), transforms.Normalize(MEAN, STD)])

    # One transform reaches every split; augmenting eval would quietly depress every metric.
    tfm = (transform or plain) if train else plain

    ds = datasets.MNIST(
        root=ds_params.get("root", "./data"), train=train, download=True, transform=tfm
    )
    return build_dataloader(ds, ds_params), {"length": len(ds), "num_classes": 10}