Skip to content

Config fields

Start from base_config(): every optional field already has a value, so a config only states what it changes.

A config, whole

What gatle-ignite init writes, minus its file docstring. It trains as it stands, so it is also the shortest correct answer to "what must a config say?". The tables below are the menu of what else it could say.

configs/my_project_v0.py
from pathlib import Path

from configs.base_utils import ckpt_dir
from gatle_ignite import base_config

IN_DIM = 64
N_CLASSES = 10


def get_config():
    cfg = base_config()  # pre-fills every optional field; override only what you need
    cfg.name = Path(__file__).stem  # -> "my_project_v0". Names the run and its checkpoints.
    cfg.project_name = "my_project"
    cfg.save_dir = ckpt_dir(cfg.name)

    cfg.main_runner = "trainer.my_project_trainer"
    cfg.model_name = "models.mlp"
    cfg.model_params = {"in_dim": IN_DIM, "hidden": 128, "n_classes": N_CLASSES}

    # Batch size lives in each split's params, not at the top level.
    common = {"in_dim": IN_DIM, "n_classes": N_CLASSES, "bs": 64, "num_workers": 0}
    cfg.train_ds_name = "dataloaders.synthetic_dataset"
    cfg.train_ds_params = {**common, "n": 2048, "seed": 0, "shuffle": True, "drop_last": True}
    cfg.valid_ds_name = "dataloaders.synthetic_dataset"
    cfg.valid_ds_params = {**common, "n": 512, "seed": 0, "shuffle": False}

    # A weighted sum even with one term, so a second term is one more entry. src_name
    # selects from the model's output dict, tgt_name from prep_batch's dict.
    cfg.criterion_name = "gatle_ignite.losses.composite"
    cfg.criterion_params = {
        "dict_of_loss_params": {
            "ce": {
                "cls_name": "losses.loss_functions.cross_entropy",
                "loss_params": {"src_name": "logits", "tgt_name": ("targets", "labels")},
                "weight": 1.0,
            }
        }
    }

    cfg.optimizer_name = "gatle_ignite.optimizers.adamw"
    cfg.optimizer_params = {"lr": 1e-3, "weight_decay": 0.01}
    cfg.lr_scheduler = "gatle_ignite.schedulers.warmup_cosine"
    cfg.lr_scheduler_params = {"warmup_epochs": 1}
    cfg.max_epochs = 15

    cfg.val_metrics = {
        "acc": {
            "cls_name": "metrics.accuracy",
            "params": {"src_name": "logits", "tgt_name": ("targets", "labels")},
        }
    }
    # score_factor = -1 for a lower-is-better score: checkpointing keeps the max.
    cfg.score_name = "valid/acc"
    cfg.score_factor = 1

    cfg.logger_name = ["text"]  # "wandb" and "discord" are also built in
    cfg.amp_dtype = "fp32"  # "bf16" once you are on a GPU that has it
    return cfg

identity / wiring

Field Type Default Notes
name (required) str unset Identifies the run: the W&B id, and the default save_dir stem.
save_dir (required) str unset Where checkpoints go. Resume reads from here too.
project_name str 'gatle' Groups runs in W&B. Not used by the text logger.
main_runner (required) str unset Dotted path to the module exposing your Trainer.
seed int 42 Seeds python, numpy and torch before anything is built.

model

Field Type Default Notes
model_name (required) str unset Dotted path to a module exposing Model(**model_params).
model_params ConfigDict {} Splatted into Model(...).
auto_model_params ConfigDict {'find_unused_parameters': False, 'sync_bn': True} kwargs for idist.auto_model (DDP wrapping). sync_bn needs a GPU backend.

data

Field Type Default Notes
train_ds_name (required) str unset Dotted path to a module exposing get_ds(ds_params, transform).
train_ds_params ConfigDict {} Passed to your get_ds. Batch size lives here (bs), not at top level.
aug_name str unset Optional. Built once and passed to EVERY split; the dataset decides who gets it.
aug_params ConfigDict {} Splatted into Transformation(...).

loss

Field Type Default Notes
criterion_name (required) str unset Usually gatle_ignite.losses.composite, even for a single term.
criterion_params ConfigDict {} For the composite: {'dict_of_loss_params': {...}}.

optimizer / scheduler

Field Type Default Notes
optimizer_name (required) str unset Dotted path. Builtins are named in full: gatle_ignite.optimizers.adamw.
optimizer_params ConfigDict {} Splatted into get_optimizer(model, ...). The LR schedule reads lr from here.
lr_scheduler str unset Dotted path. None = no scheduling.
lr_scheduler_params ConfigDict {} Splatted into get_scheduler(...).
grad_clip_norm float unset Clip gradients by total norm. None = off.
grad_clip_value float unset Clip gradients elementwise. None = off.
accum_steps int 1 Batches per optimizer step. N emulates N*bs at the memory of bs. Clipping applies to the accumulated gradient.

loop

Field Type Default Notes
max_epochs int 1 Also sets the LR schedule's geometry, so changing it on resume replays the old curve.
every_val int 1 Run the evaluator every N epochs. 0 disables validation, mirroring every_test.
every_test int 0 0 disables the tester engine entirely
train_length int unset Cap an epoch to N iterations. None = the full dataloader. Handy for smoke runs.
early_stop_patience int 0 Stop after N evaluations with no improvement in score_name. 0 = never stop early.
early_stop_after int 0 Epochs to train before the patience counter arms. 0 = arm immediately.
val_length int unset As train_length, for the evaluator.
test_length int unset As train_length, for the tester.

precision

Field Type Default Notes
amp_dtype str 'auto' auto (bf16 where supported, else fp32) | bf16 | fp16 | fp32
cudnn_benchmark bool False Autotunes per input shape. A pessimisation when shapes vary; off by default.
compile object unset torch.compile the model: True | False | "auto" (on where it is usable). Off by default; the first step pays a one-off warm-up.
compile_params ConfigDict {} Splatted into nn.Module.compile(): mode, dynamic, backend, fullgraph.

metrics / scoring

Field Type Default Notes
train_metrics ConfigDict {} {"acc": {"cls_name": "...", "params": {...}}}. Loss averages are added automatically.
val_metrics ConfigDict {} As train_metrics, on the evaluator. Keys are namespaced "valid/".
tester_metrics ConfigDict {} As train_metrics, on the tester. Keys are namespaced "test/".
score_name str unset Metric that selects the best checkpoint, e.g. "valid/acc". None = latest only.
score_factor int 1 Checkpointing keeps the MAXIMUM, so use -1 for loss/CER/WER.
tester_score_name str unset As score_name, for the test engine -> test_best_result*.
tester_score_factor int 1 As score_factor, for the test engine.

checkpointing

Field Type Default Notes
save_ckpt bool True Write checkpoints at all.
n_saved int 1 How many of each kind to keep, >= 1 or None for all. Above 1, best is by score.
resume bool False Continue from the newest latest_epoch* in save_dir. Safe to leave on; warns if there is none.
model_checkpoint_dir str '' Load WEIGHTS ONLY from this file (fine-tuning, or a clean warm start).
strict bool True strict= for the model_checkpoint_dir load.

logging

Field Type Default Notes
logger_name list ['text'] Sinks to enable. The one field taking short names: text, pbar, wandb, discord.
log_every int 100 Iterations between LR / grad-norm points.
watch_grad bool False Log gradient norms. Costs a pass over the parameters.
tags list [] Passed to W&B.

distributed

Field Type Default Notes
dist_backend str 'nccl' Only applies with >1 process; below that a run is single-process with no backend.
nnodes int 1 Number of machines. >1 needs the same command, and node_rank, on every node.
node_rank int 0 This machine's index, 0..nnodes-1. Node 0 is where master_addr must point.
nproc_per_node int unset Processes per machine. None = one per visible GPU (or 1 on CPU).
master_addr str unset Rendezvous host. None = MASTER_ADDR, else 127.0.0.1. Required for nnodes > 1.
master_port int unset Rendezvous port. None = MASTER_PORT, else random on one node. Required for nnodes > 1.

Fields not pre-set by base_config()

These are known to the framework but deliberately absent from base_config(), because their presence is what carries the meaning. A ConfigDict accepts new keys, so a config sets them directly.

Field Meaning
valid_ds_name Dotted path to the validation dataset module. Absent -> no evaluator engine.
valid_ds_params Params passed to the validation get_ds.
test_ds_name Dotted path to the test dataset module. Absent -> no tester engine.
test_ds_params Params passed to the test get_ds.
run Presence puts the trainer in inference mode: load a checkpoint and evaluate, never train.
load_from_ckpt Which checkpoint inference mode loads: "best" (default) or "latest".
wandb_entity W&B entity to log under. Absent -> your default entity.
discord_url Discord webhook URL. Prefer the DISCORD_WEBHOOK_URL env var: a webhook is a credential.

Your own fields

Any other field is left alone, so a config can carry values that only your own trainer or dataset reads.