Skip to content

Training Parameters

Parameters that control the foundation model training process. These are set in the training_params section of the YAML config or passed as a TrainingParams object in Python.

from monad.config import TrainingParams

TrainingParams

Parameter Type Default Description
epochs int 1 Number of epochs to train the model for.
learning_rate float 0.0001 Learning rate.
check_val_every_n_steps int \| None None Run validation every N training steps. Disables validation on epoch end when set.
check_val_every_n_epochs int \| None 1 Run validation every N epochs.
limit_train_batches int \| None None Limit number of training batches per epoch. Useful for quick validation of setup.
limit_val_batches int \| None None Limit number of validation batches per epoch.
loss Callable \| None None Custom loss function.
checkpoint_dir str \| Path \| None None Directory to store model checkpoints.
metric_to_monitor str \| None None Logged validation key to monitor for best-model selection and early stopping: val_loss or val_<alias>_<target_index> (for example val_auroc_0). None uses the task's default. In scenario training, fit() raises InvalidConfigurationError before training starts if the key is not logged during validation and checkpoint_dir or early_stopping is set.
metric_monitoring_mode MetricMonitoringMode \| None None Whether to minimize ("min") or maximize ("max") the monitored metric.
gradient_clip_val int \| float \| None 1.0 Gradient-norm clipping threshold: gradients are rescaled when their global norm exceeds this value. None disables clipping. In 1.13 and earlier the default was None (no clipping).
checkpoint_every_n_steps int \| None None Enable intra-epoch checkpointing at every N steps.
early_stopping EarlyStopping \| None None Early stopping configuration. See EarlyStopping below.
precision Literal[...] "bf16-mixed" / "16-mixed" Float precision for training. Defaults to "bf16-mixed" on GPUs with bfloat16 support, otherwise "16-mixed". See Precision Values below.
bf16_master_weights bool False Only with precision "bf16-true" (or "bf16"): keeps fp32 master copies of the bfloat16 model parameters inside the fused optimizer and runs the parameter update in fp32; checkpoints store the fp32 master weights. Any other precision raises a validation error (bf16_master_weights requires precision 'bf16-true'). Requires the fused Triton optimizer; if it cannot be loaded, training raises a RuntimeError.
cuda_graph bool True Captures the forward and backward passes of the training step (foundation model and scenario model) into a CUDA graph after the first few steps and replays it on later steps. Steps run eagerly instead when the model is not on a CUDA device, when the strategy is not single-device or DDP (for example FSDP), or with float16 mixed precision ("16-mixed", the default on GPUs without bfloat16 support); individual steps also run eagerly when their batch shapes differ from the captured ones (under DDP, on any rank). Set False to disable.

Inherited Parameters

These are inherited from the base parameter class and shared with TestingParams.

Parameter Type Default Description
show_entity_progress bool False Adds a progress bar that tracks processed entities next to the batch progress bars. Before each stage without a batch limit, a distinct-entity count runs on global rank 0 to give the bar its total. With the default False, no counting queries run. Has no effect when the ENABLE_PROGRESS_BAR environment variable is False.
devices list[int] \| int \| "auto" "auto" GPU devices to use. "auto" selects the least-occupied GPU automatically. It does not switch to CPU when no GPU is available: set accelerator: "cpu" for CPU runs. A positive int specifies count; a list[int] specifies device indices; -1 uses all available GPUs.
accelerator "cpu" \| "gpu" "gpu" Accelerator type.
strategy str \| None None Distributed training strategy: None (PyTorch Lightning default), "auto", "ddp", "fsdp", or "fsdp:A:B" (A = data parallelism, B = tensor parallelism).
nccl_timeout timedelta \| None None Timeout for NCCL collective operations in distributed training (DDP/FSDP). A bare number is interpreted as seconds. Ignored on a single device.
rank_sync_timeout timedelta \| None None Timeout for a dedicated per-step rank-synchronization barrier that absorbs data-loading skew between ranks, independently of nccl_timeout (which then only needs to cover the gradient sync). A bare number is interpreted as seconds. Ignored on a single device.

Note

Parameters such as metrics, top_k, predictions_threshold, entity_ids, callbacks, and approximate_decoding_params are inherited from the base class but not applicable to foundation model training. Custom metrics are explicitly rejected during pretraining. These parameters are available for Scenario Model Training.

EarlyStopping

Wraps the configuration for PyTorch Lightning's early stopping callback.

from monad.config.early_stopping import EarlyStopping

early_stopping = EarlyStopping(
    min_delta=0.001,
    patience=5,
    verbose=True,
)
Parameter Type Default Description
min_delta float 0.0 Minimum change in the monitored metric to qualify as an improvement.
patience int 3 Number of validation checks with no improvement after which training stops.
verbose bool False Whether to log information about registered improvements.

Precision Values

Valid values for the precision parameter:

Value Description
32 or "32" or "32-true" Full 32-bit float precision.
64 or "64" or "64-true" Full 64-bit float precision.
16 or "16" or "16-true" Pure 16-bit float precision.
"16-mixed" Mixed precision with float16.
"bf16" or "bf16-true" Pure bfloat16 precision.
"bf16-mixed" Mixed precision with bfloat16. Default on supported GPUs.

Tip

"bf16-mixed" is the default on GPUs that support bfloat16 (e.g., A100, H100). On older GPUs, it falls back to "16-mixed". Mixed precision provides a good balance between training speed and numerical stability.

YAML Example

In the foundation model config file:

training_params:
  learning_rate: 0.0003
  epochs: 3
  precision: "bf16-mixed"
  strategy: "ddp"
  devices: [0, 1]
  early_stopping:
    min_delta: 0.001
    patience: 5

Python Example

from monad.config import TrainingParams
from monad.config.early_stopping import EarlyStopping

training_params = TrainingParams(
    epochs=3,
    learning_rate=0.0003,
    precision="bf16-mixed",
    devices=[0, 1],
    strategy="ddp",
    early_stopping=EarlyStopping(
        min_delta=0.001,
        patience=5,
    ),
)