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 MappingModel instance.

  • train_loader – Training DataLoader.

  • val_loader – Optional validation DataLoader.

  • loss_fn – A MappingLoss or BaseLoss instance. If None, defaults to MappingLoss(RegressionLoss()).

  • config – Trainer configuration. If None, uses defaults.

  • callbacks – List of Callback instances.

  • 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.