nvalchemi-training-api

bởi nvidia

Cách cấu hình quy trình huấn luyện nvalchemi với TrainingStrategy, các hàm huấn luyện tùy chỉnh, hàm mất mát độc lập hoặc kết hợp, lịch trình trọng số mất mát, bộ tối ưu…

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.

Thêm skills từ nvidia

compileiq-debug
nvidia
Sử dụng khi có điều gì đó không ổn: Search() bị treo, tất cả các đánh giá đều trả về INVALID_SCORE, điểm số không cải thiện, mọi cấu hình đều trả về cùng một số, lỗi ptxas…
create-github-pr
nvidia
Tạo pull request GitHub bằng cách sử dụng gh CLI. Sử dụng khi người dùng muốn tạo PR mới, gửi mã để xem xét, hoặc mở pull request. Từ khóa kích hoạt -…
nemoclaw-maintainer-cross-issue-sweep
nvidia
Quét các vấn đề đang mở khác để tìm những vấn đề mà một PR nhất định có thể sửa hoặc vô tình làm hỏng. Đưa ra các cơ hội sửa lỗi liền kề và rủi ro mâu thuẫn với file:dòng…
fhir-basics
nvidia
Dạy các tác nhân cách hoạt động của API FHIR R4, những tài nguyên có sẵn, cách truy vấn chúng với tham số tìm kiếm, và cách phân tích chính xác tất cả các định dạng phản hồi…
compileiq-validate-result
nvidia
Sử dụng SAU KHI tìm kiếm hoàn tất và TRƯỚC KHI yêu cầu tăng tốc hoặc gửi ACF. Tải tệp CSV dump_results, trích xuất các ứng viên top-K (đơn mục tiêu)…
changelog-audit
nvidia
Kiểm tra Warp CHANGELOG.md trước khi phát hành: khôi phục các mục bị mất, sắp xếp theo tác động người dùng, tinh chỉnh ngôn ngữ mục, xuống dòng và (chế độ nhánh phát hành) so sánh bump…
maintain-dynamic-plugins
nvidia
Duy trì các bộ nạp plugin động NeMo Relay, tệp kê khai, SDK gốc Rust, giao thức worker gRPC, SDK worker Python, tài liệu, kiểm thử và phạm vi quy trình phát hành
dgx-diagnose
nvidia
Chẩn đoán các sự cố thường gặp của DGX Station GB300 — lỗi CUDA, nhắm sai GPU, lỗi container vLLM/SGLang, vấn đề trạng thái MIG, lỗi NVLink/Fabric Manager,…