Skip to content

Templates

The contracts table tells you a module must expose get_sampler. It does not tell you what to write. These do.

Every file below lives in examples/templates/, and every one is exercised by the test suite, so a template that stops working is caught before you copy it. Most are wired into a single config and trained together; the collate fn and the logger backend have their own tests instead, because neither fits that config: a collate reshapes the batch, and a backend is a sink rather than a component. Copy the file, rename it, point a config field at its dotted path.

You want Copy Config field Where it goes
A metric examples/templates/metric.py {train,val,tester}_metrics[<name>].cls_name metrics/
A dataset examples/templates/dataset.py train_ds_name dataloaders/
A model examples/templates/model.py model_name models/
A loss term examples/templates/loss.py criterion_params.dict_of_loss_params[<key>].cls_name losses/loss_functions/
An optimizer examples/templates/optimizer.py optimizer_name optimizer/
An LR scheduler examples/templates/scheduler.py lr_scheduler scheduler/
An augmentation examples/templates/augmentation.py aug_name augmentation/
A sampler examples/templates/sampler.py *_ds_params.sampler_params.cls_name dataloaders/data_utils/
A collate fn examples/templates/collate.py *_ds_params.collate_fn.cls_name dataloaders/data_utils/
A trainer examples/templates/trainer.py main_runner trainer/
A logger backend examples/templates/backend.py logger_name[<entry>] callbacks/backends/

There is no registry to update and no decorator to add. The dotted path in the config is the registration.

Metric

The decorators are the whole point of this file, and they are the thing most worth getting right: @sync_all_reduce on compute() is what makes the number correct under DDP. Without it the metric is right on one GPU and quietly wrong on two.

examples/templates/metric.py
"""TEMPLATE: a metric for *_metrics[<name>].cls_name, reported as "valid/<name>".

The two decorators are what keep it right under DDP.
"""

import torch
from ignite.metrics import Metric
from ignite.metrics.metric import reinit__is_reduced, sync_all_reduce

from gatle_ignite import get_value


class TopKAccuracy(Metric):
    """Fraction of samples whose true label is in the model's top k. Replace with yours."""

    def __init__(self, k=2, output_transform=lambda x: x, device="cpu"):
        self.k = k
        super().__init__(output_transform=output_transform, device=device)

    # On reset() and update(): despite the name, it clears ignite's cached result, not reduces.
    @reinit__is_reduced
    def reset(self):
        # Tensors, not python floats: sync_all_reduce moves only tensors between ranks.
        self._num_correct = torch.tensor(0, device=self._device)
        self._num_examples = torch.tensor(0, device=self._device)
        super().reset()

    @reinit__is_reduced
    def update(self, output):
        y_pred, y = output  # y_pred: (B, n_classes) logits; y: (B,) labels
        topk = y_pred.topk(min(self.k, y_pred.shape[1]), dim=1).indices
        hits = (topk == y.unsqueeze(1)).any(dim=1).sum()
        self._num_correct += hits.to(self._device)
        self._num_examples += y.shape[0]

    # All-reduces these across ranks before compute(); without it each rank scores its own
    # shard. Name EVERY attribute compute() reads: a missed one stays rank-local.
    @sync_all_reduce("_num_correct", "_num_examples")
    def compute(self):
        if self._num_examples == 0:
            # Raise, don't return 0.0: a silent zero reads as a real score.
            raise ValueError("TopKAccuracy needs at least one example")
        return self._num_correct.item() / self._num_examples.item()


def get_metric(engine_type, info, src_name="logits", tgt_name=("targets", "labels"), k=2, **kwargs):
    """engine_type is "train", "valid", "test" or a declared engine's name.

    info is your get_ds's info dict: how a metric learns num_classes without the trainer.
    """
    n_classes = (info or {}).get("num_classes")
    if n_classes is not None:
        k = min(k, n_classes)

    def output_transform(output):
        # `output` is what the step returned: {"y_pred", "target"}, plus "losses" in training.
        return get_value(output["y_pred"], src_name), get_value(output["target"], tgt_name)

    return TopKAccuracy(k=k, output_transform=output_transform, **kwargs)

Dataset

Note where the transform is applied. The framework builds one transform and hands it to every split; the dataset decides who gets it.

examples/templates/dataset.py
"""TEMPLATE: a dataset for cfg.*_ds_name.

Batch size (`bs`) lives in each split's ds_params. Never set drop_last on an eval split.
"""

import torch
from torch.utils.data import Dataset

from gatle_ignite import build_dataloader


class MyDataset(Dataset):
    def __init__(self, root=".", split="train", n=1024, seed=0, transform=None):
        self.transform = transform
        self.split = split
        # Replace with your real data. Seeded: an index must return the same sample every time.
        gen = torch.Generator().manual_seed(seed)
        self.x = torch.randn(n, 64, generator=gen)
        self.y = torch.randint(0, 10, (n,), generator=gen)

    def __len__(self):
        return len(self.x)

    def __getitem__(self, idx):
        x, y = self.x[idx], self.y[idx]
        # One transform reaches every split: apply it to train only, or you augment validation.
        if self.transform is not None and self.split == "train":
            x = self.transform(x)
        return x, y


def get_ds(ds_params, transform=None):
    """Must return (dataloader, info). Keys build_dataloader doesn't read are yours to add."""
    ds = MyDataset(
        root=ds_params.get("root", "."),
        split=ds_params.get("split", "train"),
        n=ds_params.get("n", 1024),
        seed=ds_params.get("seed", 0),
        transform=transform,
    )

    # build_dataloader, not DataLoader: it shards across ranks and splits bs by world size.
    dataloader = build_dataloader(ds, ds_params)

    # "length" is required. Other keys reach every metric as its `info` (num_classes, say).
    info = {"length": len(ds), "num_classes": 10}
    return dataloader, info

Model

examples/templates/model.py
"""TEMPLATE: a model for cfg.model_name; model_params are splatted into Model(**model_params)."""

from torch import nn


class Model(nn.Module):
    """The class name must be exactly `Model`: that is what the framework imports."""

    def __init__(self, in_dim=64, hidden=128, n_classes=10):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(in_dim, hidden),
            nn.ReLU(),
            nn.Linear(hidden, n_classes),
        )

    def forward(self, x):
        """Return a DICT, not a tensor: a config's `src_name` selects from its keys.

        prep_batch's "model_input" is splatted in, so {"x": ...} needs a parameter named `x`.
        """
        return {"logits": self.net(x)}

Loss

A sub-loss returns a bare tensor. Only the composite returns (total, {"loss_<key>": tensor}).

examples/templates/loss.py
"""TEMPLATE: a loss term, one entry in criterion_params.dict_of_loss_params.

Each term is charted as its WEIGHTED contribution, as `train/loss_<key>_avg`.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F

from gatle_ignite import get_value


class Loss(nn.Module):
    """The class name must be exactly `Loss`.

    Return a BARE TENSOR: only the composite returns (total, {"loss_<key>": tensor}).
    """

    def __init__(self, src_name="logits", tgt_name=("targets", "labels"), gamma=2.0):
        super().__init__()
        # src_name/tgt_name let one loss module serve a different head without editing it.
        self.src_name = src_name
        self.tgt_name = tgt_name
        self.gamma = gamma

    def forward(self, y_pred, target, iteration=None):
        """y_pred is the model's output dict; target is prep_batch's whole return.

        `iteration` is the global step, for a term that ramps in; ignore it otherwise.
        """
        logits = get_value(y_pred, self.src_name)  # (B, C)
        labels = get_value(target, self.tgt_name)  # (B,)
        # A shape mismatch that broadcasts does not raise: it trains flat. Check both shapes.

        # Focal loss. Replace with whatever your term actually is.
        ce = F.cross_entropy(logits, labels, reduction="none")
        p_t = torch.exp(-ce)
        return ((1 - p_t) ** self.gamma * ce).mean()

Optimizer

examples/templates/optimizer.py
"""TEMPLATE: an optimizer for cfg.optimizer_name; optimizer_params are splatted in as kwargs."""

import torch


def get_optimizer(model, lr=1e-3, weight_decay=0.0, no_decay_on_norm_and_bias=True, **kwargs):
    """Return the optimizer object the trainer holds; it need not be a torch.optim.Optimizer.

    `model` is already DDP-wrapped: parameters() sees through it, but a submodule does not.
    Reach one via `getattr(model, "module", model).gen`, or it fails only on the cluster.
    The framework asks only for param_groups, state_dict, load_state_dict, zero_grad, step.
    """
    if not no_decay_on_norm_and_bias:
        return torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay, **kwargs)

    # Biases and norm scales have too few parameters to overfit; decaying them mostly hurts.
    decay, no_decay = [], []
    for name, param in model.named_parameters():
        if not param.requires_grad:
            continue
        if param.ndim <= 1 or name.endswith(".bias"):
            no_decay.append(param)
        else:
            decay.append(param)

    return torch.optim.AdamW(
        [
            {"params": decay, "weight_decay": weight_decay},
            {"params": no_decay, "weight_decay": 0.0},
        ],
        lr=lr,
        **kwargs,
    )

Scheduler

The 3-tuple return is the contract: what to attach, to which engine, on which event.

examples/templates/scheduler.py
"""TEMPLATE: an LR scheduler for cfg.lr_scheduler; lr_scheduler_params are splatted in."""

from ignite.engine import Events
from ignite.handlers import LRScheduler
from torch.optim.lr_scheduler import StepLR


def get_scheduler(cfg, train_dl, optimizer, engines, step_every_epochs=5, gamma=0.1):
    """Return (scheduler, engine_name, event); a list of schedulers works too.

    engine_name is any key of `engines`, whose value is None for an engine not built. An
    evaluator-driven scheduler is how a plateau schedule sees validation metrics.
    """
    # Wrap a torch scheduler in ignite's LRScheduler: it is attached as an event handler, so it
    # must be callable. Its first call applies the initial value, so N epochs give N-1 decays.
    scheduler = LRScheduler(StepLR(optimizer, step_size=step_every_epochs, gamma=gamma))
    return scheduler, "trainer", Events.EPOCH_COMPLETED

Augmentation

examples/templates/augmentation.py
"""TEMPLATE: an augmentation for cfg.aug_name; aug_params are splatted into Transformation."""

import torch


class Transformation:
    """The class name must be exactly `Transformation`. Any callable will do.

    It is built once and handed to EVERY split: your dataset decides which ones apply it.
    """

    def __init__(self, p=0.5, scale=0.1):
        self.p = p
        self.scale = scale

    def __call__(self, x):
        if torch.rand(1).item() < self.p:
            x = x + torch.randn_like(x) * self.scale
        return x

Sampler

examples/templates/sampler.py
"""TEMPLATE: a sampler for *_ds_params.sampler_params.cls_name. Drop `shuffle` when you add one."""

import torch
from torch.utils.data import WeightedRandomSampler


def get_sampler(dataset, num_samples=None, **kwargs):
    """Return a single-process torch Sampler; under DDP the framework shards it per rank.

    Read labels OFF the dataset (a `labels` property), never by iterating it: that runs
    __getitem__, augmentation and decoding included, per sample. The fallback is slow.
    """
    if hasattr(dataset, "labels"):
        labels = torch.as_tensor([int(y) for y in dataset.labels])
    else:
        labels = torch.as_tensor([int(y) for _, y in dataset])
    counts = torch.bincount(labels).clamp(min=1)
    weights = (1.0 / counts.float())[labels]

    return WeightedRandomSampler(
        weights=weights,
        num_samples=num_samples or len(dataset),
        replacement=True,
        **kwargs,
    )

Collate

examples/templates/collate.py
"""TEMPLATE: a collate fn for *_ds_params.collate_fn.cls_name; params go to get_collate_fn.

ds_params["collate_fn"] also takes a plain callable; the dotted form lets a config choose.
"""

import torch


def get_collate_fn(pad_id=0):
    """Return the callable torch hands a list of samples.

    It must be picklable for num_workers > 0: a closure here is, a lambda in your get_ds is not.
    """

    def collate(batch):
        # Any shape will do: prep_batch maps whatever this returns into model_input/targets.
        xs, ys = zip(*batch)
        lengths = torch.tensor([len(x) for x in xs])
        padded = torch.nn.utils.rnn.pad_sequence(xs, batch_first=True, padding_value=pad_id)
        return {"x": padded, "lengths": lengths, "y": torch.stack(ys)}

    return collate

Trainer

In the normal case this is one prep_batch and nothing else. The rest is there to show you what the hooks look like when you do need them. Delete what you don't.

examples/templates/trainer.py
"""TEMPLATE: a trainer for cfg.main_runner. Usually just prep_batch: delete what you don't need."""

import torch

from gatle_ignite import BaseTrainer, EngineSpec, to_device


class Trainer(BaseTrainer):
    """The class name must be exactly `Trainer`."""

    def prep_batch(self, batch, split="train", **kwargs):
        """Map a raw batch into two dicts.

        `model_input` is splatted into the model's forward, so its keys are that signature.
        `targets` is what a config's tgt_name walks. `split` is which engine is asking.
        Keep the **kwargs, but don't read from it.
        """
        x, y = batch
        return to_device({"model_input": {"x": x}, "targets": {"labels": y}})

    # ---- Everything below is optional. Delete it unless you need it. ----

    def eval_specs(self):
        """Add an engine: its loader, metrics, cadence, logging and checkpoint score follow.

        Needs cfg.probe_ds_name, cfg.probe_metrics, cfg.every_probe.
        """
        return super().eval_specs() + (EngineSpec.for_split("probe"),)

    def backward(self, loss, step=True):
        """Custom gradient handling. This hook owns optimizer.step(), clipping and the scaler.

        Honour `step`: it is False on every batch but the last of an accumulation window.
        `loss` is already divided by the window, so do not rescale it.
        """
        loss.backward()
        for name, param in self.model.named_parameters():
            if param.grad is not None and "backbone" in name:
                param.grad *= 0.1  # e.g. a smaller effective LR for pretrained layers
        if step:
            self._clip_grads(self.model)
            self.optimizer.step()

    def forward(self, model_input):
        """Override when the call isn't a plain splat: extra args, a two-pass model."""
        return self.model(**model_input)

    def eval_step(self, engine, batch, split="valid"):
        """Override when eval is NOT a forward pass (sampling, a decode). Keep all four steps."""
        # 1. Release the last batch and output, or an epoch's worth stays alive.
        engine.state.batch = None
        engine.state.output = None
        self.model.eval()

        # 2. prep_batch is yours to call. The engine does not call it for you.
        x = self.prep_batch(batch, split=split)

        # 3. Re-enter autocast yourself, with the dtype the config resolved.
        with torch.no_grad():
            with torch.autocast(
                device_type=self.device_type, dtype=self.dtype, enabled=self.autocast_enabled
            ):
                y_pred = self.forward(x["model_input"])

        # 4. "target" is prep_batch's WHOLE return: a tgt_name is a path into it.
        return {"y_pred": y_pred, "target": x}

    def train_step(self, engine, batch, split="train"):
        """Override when the step isn't standard supervised, such as a GAN's alternating updates.

        Keep eval_step's steps. `losses` MUST carry "loss" and a `loss_<key>` per crit_keys
        entry, or the first epoch ends in a KeyError. An override drops gradient accumulation.
        """
        engine.state.batch = None
        engine.state.output = None
        self.model.train()
        self.optimizer.zero_grad(set_to_none=True)

        x = self.prep_batch(batch, split=split)
        with torch.autocast(
            device_type=self.device_type, dtype=self.dtype, enabled=self.autocast_enabled
        ):
            y_pred = self.forward(x["model_input"])
            loss, dict_losses = self.loss_fn(y_pred, x, iteration=engine.state.iteration)

        # Backward outside autocast, which is for the forward pass only.
        self.backward(loss)
        return {"y_pred": y_pred, "target": x, "losses": {"loss": loss, **dict_losses}}

Backend

logger_name is the one field that is not a plain dotted path: it is a list of enabled sinks, so a short name resolves to a builtin and anything else is a path to a module like this. Both kinds run at once. Subclass the base: it supplies no-op watch/finish, and the framework calls both without checking they exist.

examples/templates/backend.py
"""TEMPLATE: a logging sink, enabled by a dotted entry in cfg.logger_name. This one writes a CSV."""

import csv
from pathlib import Path

from gatle_ignite.callbacks.logging import Backend as BaseBackend


class Backend(BaseBackend):
    """The name `Backend` is what the framework imports.

    Subclass the base: the framework calls `watch` and `finish` unguarded, and it has both.
    """

    def __init__(self, cfg):
        super().__init__(cfg)
        # Third-party imports go in here, not at module top, so only runs using this sink pay.
        self.path = Path(cfg.save_dir) / f"{cfg.name}_metrics.csv"
        self.path.parent.mkdir(parents=True, exist_ok=True)
        self._columns = None

    def log(self, metrics, step=None, epoch=None):
        """`metrics` is flat {str: float}, keyed by engine ("valid/top2").

        `step` or `epoch` may be None, depending on which event fired: never assume both.
        """
        row = {"step": step, "epoch": epoch, **metrics}
        # Keys vary by event, so the first row fixes the header and later extra keys are dropped.
        write_header = self._columns is None
        if write_header:
            self._columns = list(row)
        with self.path.open("a", newline="") as fh:
            writer = csv.DictWriter(fh, fieldnames=self._columns, extrasaction="ignore")
            if write_header:
                writer.writeheader()
            writer.writerow(row)

    def watch(self, model):
        """Called once per run, only when cfg.watch_grad is set. Usually a no-op."""

    def finish(self, failed=False):
        """Called at the end of EVERY run. `failed` is True if training raised."""

See prep_batch for the one hook most tasks write, and eval_specs for adding an engine.