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.
"""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.
"""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¶
"""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}).
"""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¶
"""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.
"""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¶
"""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¶
"""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¶
"""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.
"""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.
"""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.