3. prep_batch¶
trainer/my_trainer.py → cfg.main_runner
This is normally the whole task-specific trainer. Engines, metrics, checkpointing, resume, logging, AMP and DDP are inherited.
from gatle_ignite import BaseTrainer, to_device
class Trainer(BaseTrainer):
def prep_batch(self, batch, split="train", **kwargs):
x, y = batch
return to_device({"model_input": {"x": x}, "targets": {"labels": y}})
model_inputis splatted into the model'sforward, so its keys match that signature.targetsis what the loss and metrics resolve theirtgt_namepaths against.
split is the engine's engine_type: "train", "valid", "test", or whatever an added engine
calls itself. Read split, never kwargs: a mistyped kwargs key returns None and silently takes
the wrong path.
The launcher constructs Trainer(local_rank, cfg) and drives the lifecycle, so you never call
fit() yourself.
train_step and eval_step are the hooks below this one: override eval_step when eval is not a
forward pass (a diffusion sampler, an autoregressive decode), and train_step when the step itself
is not standard supervised. Both have working defaults, as do build_model, build_dataloaders,
forward, backward, eval_specs, train_spec, eval_context and
extra_to_save, which are documented in the Python API.
The loss receives prep_batch's whole return
Not x["targets"]. Every contract in this framework passes (y_pred, target) where target is
what you returned here, which is why tgt_name paths start ("targets", ...).
Templates: trainer