mapping_networks.trainers.MappingTrainer¶
- class mapping_networks.trainers.MappingTrainer(model: MappingModel, train_loader: DataLoader[Any], val_loader: DataLoader[Any] | None = None, loss_fn: MappingLoss | BaseLoss | None = None, config: TrainerConfig | None = None, callbacks: list[Callback] | None = None, batch_adapter: BatchAdapter | None = None)¶
Training loop for mapping networks.
The trainer owns optimization only: forward, compute loss, backward, optimizer step, validation, and callbacks. Mapping logic stays in the model/generator/loss subsystems.
- Parameters:
model – A
MappingModelinstance.train_loader – Training
DataLoader.val_loader – Optional validation
DataLoader.loss_fn – A
MappingLossorBaseLossinstance. IfNone, defaults toMappingLoss(RegressionLoss()).config – Trainer configuration. If
None, uses defaults.callbacks – List of
Callbackinstances.batch_adapter – Adapter for unpacking dataloader batches. Defaults to
TupleBatchAdapter.
Example:
trainer = MappingTrainer( model=mapping_model, train_loader=train_loader, loss_fn=MappingLoss(ClassificationLoss()), config=TrainerConfig(max_epochs=50, learning_rate=3e-4), callbacks=[EarlyStopping(patience=5), MetricLogger()], ) history = trainer.fit()
- __init__(model: MappingModel, train_loader: DataLoader[Any], val_loader: DataLoader[Any] | None = None, loss_fn: MappingLoss | BaseLoss | None = None, config: TrainerConfig | None = None, callbacks: list[Callback] | None = None, batch_adapter: BatchAdapter | None = None) None¶
Methods
__init__(model, train_loader[, val_loader, ...])fit()Run the full training loop.