Trainer and callbacks¶
Status: Implemented in Milestone 8.
The MappingTrainer class manages the training loop, validation loop, device placement, precision scaling, and callback notifications. It keeps the high-level mapping and loss logic clean by managing only the optimization process.
Quick start¶
import torch
from torch.utils.data import DataLoader, TensorDataset
from mapping_networks import MappingModel, MappingTrainer, TrainerConfig
from mapping_networks.callbacks import EarlyStopping, MetricLogger
# Create model and loaders
model = MappingModel(target_model, latent_dim=128)
train_loader = DataLoader(TensorDataset(X_train, y_train), batch_size=32)
val_loader = DataLoader(TensorDataset(X_val, y_val), batch_size=32)
# Configure trainer
config = TrainerConfig(
max_epochs=50,
learning_rate=1e-3,
optimizer="adamw",
scheduler="cosine",
device="auto",
seed=42,
)
trainer = MappingTrainer(
model=model,
train_loader=train_loader,
val_loader=val_loader,
config=config,
callbacks=[
EarlyStopping(monitor="val_loss", patience=5),
MetricLogger(log_every_n_batches=10),
],
)
history = trainer.fit()
Batch adapters¶
DataLoader output formats vary across datasets. To support arbitrary formats without modifying the training loop, MappingTrainer delegates batch unpacking to a BatchAdapter.
Built-in adapters¶
TupleBatchAdapter¶
Unpacks standard PyTorch datasets returning (inputs, targets) or (inputs, targets, *extra). Extra elements are ignored.
from mapping_networks.trainers import TupleBatchAdapter
# Automatically used by default in MappingTrainer
adapter = TupleBatchAdapter()
MappingBatchAdapter¶
Unpacks dict-like batches (common in Hugging Face or text datasets) by key:
from mapping_networks.trainers import MappingBatchAdapter
# For a batch returning {"pixel_values": X, "label": y}
adapter = MappingBatchAdapter(input_key="pixel_values", target_key="label")
Custom batch adapter¶
If your data format is complex (e.g. structured keywords), subclass BatchAdapter and override unpack:
from typing import Any
from mapping_networks.trainers import BatchAdapter
class CustomBatchAdapter(BatchAdapter):
def unpack(self, batch: Any) -> tuple[tuple[Any, ...], Any]:
# Return positional arguments to pass to target forward, and targets
return (batch.features, batch.metadata), batch.label
Training features¶
Device selection: Resolves
"auto"to CUDA when available, otherwise CPU. Automatically places the model, data, and PyTorch modules on the target device.AMP (Automatic Mixed Precision): Enabled with
amp_enabled=TrueinTrainerConfig. Usestorch.amp.autocastandGradScalerfor training.Gradient accumulation: Splits a batch size logically across multiple forward steps to save memory. Set
accumulation_steps > 1.Gradient clipping: Supports clipping gradients either by norm or absolute value:
gradient_clip_norm: floatgradient_clip_value: float
Callbacks¶
Subclass Callback to observe or modify the training run at key life cycle hooks:
from mapping_networks import Callback
class CustomCallback(Callback):
def on_fit_start(self, trainer):
print("Training starting!")
def on_epoch_end(self, trainer, epoch, metrics):
print(f"Epoch {epoch} finished. Metrics: {metrics}")
Hooks lifecycle¶
on_fit_start(trainer)For each epoch: a.
on_epoch_start(trainer, epoch)b. For each batch:on_batch_start(trainer, batch_idx)→on_batch_end(trainer, batch_idx, loss_output)c.on_validation_start(trainer)→on_validation_end(trainer, metrics)d.on_epoch_end(trainer, epoch, metrics)on_fit_end(trainer)
Callbacks can alter trainer behavior (e.g., setting trainer.should_stop = True to terminate the training loop).
Built-in callbacks¶
EarlyStopping(monitor, patience, min_delta, mode): Stops training early if a metric (e.g.val_loss) stops improving.MetricLogger(log_every_n_batches): Prints epoch and step loss metrics using standard Python logging.