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_augtest_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}