nvalchemi-training-api

par nvidia

Comment configurer les workflows d’entraînement nvalchemi avec TrainingStrategy, des fonctions d’entraînement personnalisées, des pertes autonomes ou composées, des programmes de pondération des pertes, un optimiseur…

npx skills add https://github.com/nvidia/nvalchemi-toolkit --skill nvalchemi-training-api

nvalchemi Training API

Overview

Use TrainingStrategy as the owner of one training job: model(s), dataloaders, loss, optimizer/scheduler config, validation, hooks, runtime counters, and checkpoints. For full details, see docs/userguide/training.md, docs/userguide/losses.md, and docs/modules/training/checkpoints.rst.

import torch

from nvalchemi.data import Batch
from nvalchemi.models.base import BaseModelMixin
from nvalchemi.training import (
    CheckpointHook,
    ComposedLossFunction,
    CosineWeight,
    EnergyMSELoss,
    ForceMSELoss,
    LinearWeight,
    OptimizerConfig,
    StressMSELoss,
    TrainingStrategy,
    ValidationConfig,
    create_model_spec,
)

Minimal Pattern

loss_fn = ComposedLossFunction(
    [EnergyMSELoss(), ForceMSELoss()],
    weights=[1.0, 10.0],
    normalize_weights=False,
)

strategy = TrainingStrategy(
    models=model,
    optimizer_configs=OptimizerConfig(
        optimizer_cls=torch.optim.AdamW,
        optimizer_kwargs={"lr": 1e-4, "weight_decay": 1e-5},
    ),
    loss_fn=loss_fn,
    validation_config=ValidationConfig(validation_data=val_loader, every_n_epochs=1),
    hooks=[CheckpointHook("runs/example/checkpoints", epoch_interval=1)],
    num_epochs=20,
)
strategy.run(train_loader)

Model-Agnostic Inputs

Accept any torch.nn.Module that works with the selected training_fn. Prefer wrapped BaseModelMixin models for standard AtomicData/Batch contracts; see the nvalchemi-model-wrapping skill or docs/userguide/models.md when adapting arbitrary MLIPs.

Make model construction reproducible when possible. Use native checkpoint constructors that carry a spec, or store a create_model_spec(...) for custom wrappers so strategy checkpoints can rebuild the model before loading weights. Treat foreign checkpoints as imported weights until a fresh TrainingStrategy checkpoint has been saved.


Custom Training Functions

Use training_fn when the batch needs custom routing, multiple models, teacher outputs, auxiliary predictions, or non-standard model outputs. It receives (model, batch) for a single model or (models, batch) for named models and returns the prediction mapping consumed by loss_fn.

For multiple models, pass a named mapping. optimizer_configs must use the same model keys for trainable models. Models absent from optimizer_configs may be used in the forward path but are frozen during training.

def training_fn(models: dict[str, BaseModelMixin], batch: Batch):
    student = models["student"](batch)
    with torch.no_grad():
        teacher = models["teacher"](batch)
    return {
        "student_energy": student["energy"],
        "teacher_energy": teacher["energy"].detach(),
    }

loss_fn = ComposedLossFunction(
    [EnergyMSELoss(prediction_key="student_energy", target_key="teacher_energy")]
)

strategy = TrainingStrategy(
    models={"student": student_model, "teacher": teacher_model},
    optimizer_configs={
        "student": [
            OptimizerConfig(
                optimizer_cls=torch.optim.AdamW,
                optimizer_kwargs={"lr": 3e-5},
            )
        ]
    },
    training_fn=training_fn,
    loss_fn=loss_fn,
    num_epochs=5,
)

If targets do not come directly from the batch, also provide a loss_target_assembler; see docs/userguide/training.md.


Losses And Scheduling

A standalone leaf loss such as EnergyMSELoss() can be used when the objective has one target. Use ComposedLossFunction or operator sugar for multi-target objectives. Leaf losses consume unweighted tensors; weights and schedules live on the composition. Built-in schedules include ConstantWeight, LinearWeight, CosineWeight, and PiecewiseWeight.

Built-in losses default to dtype_policy="strict" and raise when prediction and target dtypes differ. When building or reviewing workflows, check likely label/model dtype alignment, such as float64 dataset labels with float32 model outputs. If the mismatch is intentional, tell the user they can set dtype_policy="prediction_to_target" to cast outputs to labels or dtype_policy="target_to_prediction" to cast labels to outputs. Set the policy on an explicit ComposedLossFunction(...), on a leaf loss, or after operator-sugar construction:

loss_fn = EnergyMSELoss() + ForceMSELoss()
loss_fn.dtype_policy = "prediction_to_target"

A leaf loss with its own explicit dtype_policy overrides the composed-level policy. The setting is included in serializable loss specs for restartable training workflows. For CLI scaffolds, pass --loss-dtype-policy strict, --loss-dtype-policy prediction_to_target, or --loss-dtype-policy target_to_prediction to nvalchemi-training train init or nvalchemi-training finetune init ...; spec report shows the selected policy.

loss_fn = (
    1.0 * EnergyMSELoss()
    + LinearWeight(start=0.0, end=10.0, num_steps=1000) * ForceMSELoss()
    + CosineWeight(start=0.0, end=0.1, num_steps=5000) * StressMSELoss()
)

Caveats:

  • normalize_weights=True is the default; set False for raw coefficient sums.
  • per_epoch=True schedules require epoch during loss calls.
  • Custom schedules must implement per_epoch, __call__(step, epoch), and to_spec() if they are used in restartable strategy checkpoints.
  • For custom leaf-loss internals, use nvalchemi-loss-api and docs/userguide/losses.md.

Optimizers And Schedulers

Use OptimizerConfig(optimizer_cls=..., optimizer_kwargs=...); add scheduler_cls and scheduler_kwargs when needed. Keyword arguments are validated against class constructors before training starts.

Time-based schedulers step after optimizer steps. ReduceLROnPlateau-style metric schedulers step after validation; set scheduler_metric_adapter to a validation-summary key or callable when the default "total_loss" is not right.


Checkpoints And Reproducibility

Training workflows should be fully checkpointable and reproducible:

  • Use deterministic model/wrapper constructors or create_model_spec(...).
  • Keep loss functions, schedules, optimizer configs, and restart-critical hooks serializable; implement to_spec() where protocols require it.
  • Use CheckpointHook for periodic checkpoints and save early enough for preempted jobs, including Slurm-style cluster runs.
  • Make data splits, sampler state, seeds, units, dtype/device choices, and config files explicit in the run directory.
  • For multi-GPU or multi-node runs (DDP, rank-safe checkpointing), see the Scaling to multiple GPUs section below.

Strategy checkpoints are restart packages: model weights, optimizer and scheduler state, strategy counters, checkpointable hook state, and reconstruction metadata.


Resume Training

Use resume when continuing the same run after interruption. This is different from fine-tuning, which imports weights into a new objective or dataset.

strategy = TrainingStrategy.load_checkpoint("runs/example/checkpoints", map_location="cuda")
strategy.run(train_loader)

Resume only from native TrainingStrategy checkpoints when optimizer, scheduler, hook state, and counters matter. Plain pretrained weight files are not sufficient for faithful continuation. To start a fresh fine-tuning run from native checkpoint weights, use FineTuningStrategy.from_pretrained_checkpoint(...) from nvalchemi-fine-tuning; opt into source loss or optimizer classes with use_original_loss=True or use_original_opt_class=True when those defaults are desired. See docs/modules/training/checkpoints.rst.


Scaling to multiple GPUs (DDP)

Data-parallel training routes through DistributedManager (re-exported from PhysicsNeMo as nvalchemi.distributed.DistributedManager); prefer it as the single entry point. It owns rank, device, and process-group state, and passing it to TrainingStrategy alongside a DDPHook gives every hook the same runtime view, so one script runs unchanged on one process or many (with world size one, DDPHook is a no-op). See docs/userguide/distributed_training.md for the full guide.

from nvalchemi.distributed import DistributedManager
from nvalchemi.training.hooks import DDPHook

DistributedManager.initialize()          # also handles single-process runs
manager = DistributedManager()

strategy = TrainingStrategy(
    models=model,
    optimizer_configs=OptimizerConfig(
        optimizer_cls=torch.optim.AdamW, optimizer_kwargs={"lr": 1e-4}
    ),
    loss_fn=EnergyMSELoss() + ForceMSELoss(),
    distributed_manager=manager,
    hooks=[DDPHook(), CheckpointHook("runs/ddp/checkpoints", epoch_interval=1)],
    num_epochs=20,
)
strategy.run(train_loader)

DDPHook (during strategy setup) wraps the trainable models in DistributedDataParallel, selects the rank-local device, and injects a distributed sampler into the active dataloader, so no manual sampler wiring is needed. Rank-safety is handled for you: validation all-reduces its metrics across ranks (so never rank-gate the validation call), and CheckpointHook writes from global rank 0 only, unwrapping DDP so checkpoints store plain weights. Reporting is rank-aware too (see nvalchemi-reporting).

Launch one process per GPU with torchrun:

torchrun --standalone --nproc_per_node=4 train.py
# runnable example:
uv run --extra cuXX torchrun --standalone --nproc_per_node=2 \
    examples/intermediate/06_ddp_mlp_training.py --backend auto

For multi-node launches (torchrun --nnodes/--rdzv_endpoint or Slurm srun), the rank helpers (get_rank, get_world_size, barrier, all_reduce in nvalchemi/training/distributed.py), and sampler/backend tuning, see docs/userguide/distributed_training.md.

Plus de skills de nvidia

compileiq-debug
nvidia
Utilisez quand quelque chose ne va pas : Search() bloque, toutes les évaluations retournent INVALID_SCORE, les scores ne s'améliorent pas, chaque configuration retourne le même nombre, erreurs ptxas…
create-github-pr
nvidia
Créer des pull requests GitHub en utilisant l'interface en ligne de commande gh. Utiliser lorsque l'utilisateur souhaite créer une nouvelle PR, soumettre du code pour révision, ou ouvrir une pull request. Mots-clés de déclenchement -…
nemoclaw-maintainer-cross-issue-sweep
nvidia
Analyse les autres problèmes ouverts pour trouver ceux qu’une PR donnée pourrait également corriger ou casser accidentellement. Génère des opportunités de correctifs adjacents et des risques de contradiction avec fichier:ligne…
fhir-basics
nvidia
Apprend aux agents comment fonctionnent les API FHIR R4, quelles ressources sont disponibles, comment les interroger avec des paramètres de recherche, et comment analyser correctement tous les formats de réponse…
compileiq-validate-result
nvidia
Utiliser APRÈS qu'une recherche soit terminée et AVANT de réclamer un accélérateur ou d'expédier un ACF. Charge le CSV dump_results, extrait les K meilleurs candidats (mono-objectif)…
changelog-audit
nvidia
Auditer le CHANGELOG.md de Warp avant une publication : récupérer les entrées perdues, trier par impact utilisateur, affiner le langage des entrées, ajuster les retours à la ligne et (en mode branche de publication) mettre à jour la comparaison…
maintain-dynamic-plugins
nvidia
Maintenir les chargeurs de plugins dynamiques NeMo Relay, les manifestes, les SDK natifs Rust, le protocole worker gRPC, le SDK worker Python, la documentation, les tests et la couverture du workflow de publication
dgx-diagnose
nvidia
Diagnostiquer les problèmes courants du DGX Station GB300 — plantages CUDA, ciblage incorrect du GPU, bugs de conteneur vLLM/SGLang, problèmes d'état MIG, erreurs NVLink/Fabric Manager,…