Custom Loss & Callbacks
Custom Loss Function
For imbalanced binary classification, TrainingParams has no built-in class weights or class-balanced oversampling option; pass a custom weighted loss such as weighted BCE through loss, and monitor AveragePrecision with AUROC for the rare positive class. The pos_weight value is data-dependent, so calculate it from the training split rather than copying the example value.
The default loss functions are optimized for each task type. You can replace them with any loss function from PyTorch or written in Python.
Using a PyTorch Loss
from torch.nn.functional import mse_loss
from monad.config import TrainingParams
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=1,
loss=mse_loss,
)
Defining a Custom Loss
Monad calls the loss function as loss(output, target, weight=weights). Accept the weight argument and apply it before averaging: it carries the per-example loss weights and zeroes out padding rows, so a loss that ignores it trains on a different objective. Create any constant tensor from the input, as input.new_tensor(...) does below, so it lands on whichever device holds the batch.
from torch.nn.functional import binary_cross_entropy_with_logits
from monad.config import TrainingParams
def weighted_bce(input, target, weight=None, size_average=None,
reduce=None, reduction="mean"):
return binary_cross_entropy_with_logits(
input, target, weight, size_average, reduce, reduction,
pos_weight=input.new_tensor([0.9]),
)
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=1,
loss=weighted_bce,
)
Weighted Multiclass Cross-Entropy
In Monad 1.14, MulticlassClassificationTask and TrainingParams do not expose a one-line class-weight argument. Pass weighted cross-entropy through TrainingParams.loss. Multiclass targets are normalized class vectors in the same fixed order as class_names; PyTorch cross-entropy accepts these probability vectors directly, so do not convert them with an order-changing label encoder.
Calculate the weights from the training split only. Persist the counts in class-name order so the same function and weights can be recreated before every load_from_checkpoint call. For example, write training_class_counts.json next to the script as a JSON list of positive counts, one per class:
import json
from pathlib import Path
import torch
from torch.nn.functional import cross_entropy
from monad.config import TrainingParams
def inverse_frequency_weights(class_counts: list[int]) -> torch.Tensor:
counts = torch.tensor(class_counts, dtype=torch.float32)
if counts.ndim != 1 or len(counts) < 2 or torch.any(counts <= 0):
raise ValueError("class counts must contain one positive training count per class")
return counts.sum() / (len(counts) * counts)
# The JSON order must exactly match MulticlassClassificationTask(class_names=CLASS_NAMES).
counts_path = Path(__file__).with_name("training_class_counts.json")
training_class_counts = json.loads(counts_path.read_text())
CLASS_WEIGHTS = inverse_frequency_weights(training_class_counts)
def weighted_multiclass_cross_entropy(logits, target, weight=None):
loss = cross_entropy(
logits,
target.to(device=logits.device, dtype=logits.dtype),
weight=CLASS_WEIGHTS.to(device=logits.device, dtype=logits.dtype),
reduction="none",
).unsqueeze(1)
if weight is not None:
loss = loss * weight
return loss.mean()
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=1,
loss=weighted_multiclass_cross_entropy,
)
Treat weighting as a validation-controlled comparison, not as an automatic fix. Keep the target, class order, split, and cohort fixed; compare against the default loss on validation macro-F1 and per-class recall; keep the final test untouched. A zero training count cannot receive a finite inverse weight — fix the training population or class definition rather than hiding it with a constant.
Redefine custom loss in all load scripts
If you use a custom loss function, define it in every script that loads the trained model via load_from_checkpoint. The loss function must be importable at load time.
Callbacks
Attach PyTorch Lightning callbacks to supplement training with additional functionality — progress bars, custom logging, learning-rate scheduling, etc.
from pytorch_lightning.callbacks import TQDMProgressBar
from monad.config import TrainingParams
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=1,
callbacks=[TQDMProgressBar(refresh_rate=100)],
)
Pass callbacks as a list via TrainingParams.callbacks. Any callback compatible with the PyTorch Lightning Trainer can be used.
Early Stopping
Early stopping halts training when the monitored metric stops improving, preventing overfitting and saving compute time. Configure it via TrainingParams.early_stopping:
from monad.config.early_stopping import EarlyStopping
from monad.config import TrainingParams, MetricMonitoringMode
training_params = TrainingParams(
checkpoint_dir="./model",
epochs=20,
metric_to_monitor="val_auroc_0",
metric_monitoring_mode=MetricMonitoringMode.MAX,
early_stopping=EarlyStopping(patience=5),
)
| Parameter | Default | Description |
|---|---|---|
patience | 3 | Number of validation checks with no improvement before stopping. Must be greater than 1. |
min_delta | 0.0 | Minimum change to qualify as an improvement |
verbose | False | Log a message when early stopping triggers |
Use early stopping to prevent overfitting
Without early stopping, the best checkpoint is still saved based on metric_to_monitor, but training runs for the full epochs count. Early stopping is most useful when you set a high epochs ceiling and want training to finish as soon as gains plateau.