nemo_automodel.recipes.llm.train_ft
nemo_automodel.recipes.llm.train_ft
Module Contents
Classes
Functions
Data
API
Bases: BaseRecipe
Recipe for fine-tuning a model for next-token prediction.
This class orchestrates training, from setup to main training loop.
Broadcast a PP last-stage scalar to the other ranks in its pipeline group.
Collect MoE load balance metrics with DP all-reduce.
Must be called on ALL ranks (the all-reduce is collective).
Stores the result in self._moe_layer_loads for rank-0 logging.
Configure every local model stage and return its NEAT data requirements.
Create the distributed setup used by this recipe rank.
Run one local batch and accumulate its loss and optional gradients.
Parameters:
Microbatch index in the accumulation window.
Input mapping with token IDs, labels, and physical NEAT document IDs of shape [batch, sequence]. NEAT attention metadata is batch-major; legacy THD inputs are flattened by the sharder. THD MTP requires physical boundaries in cu_seqlens_padded of shape [num_sequences + 1] or [1, num_sequences + 1], or _packed_seq_ids [batch, sequence] supplied by the native model sharder from physical seq_lens_padded [batch, num_sequences].
List receiving the detached scalar loss.
Global supervised-token count for loss normalization.
Number of microbatches in the accumulation window.
Whether to backpropagate the combined main and MTP loss.
Log MoE load balance metrics to wandb.
Call after _collect_moe_load_balance. Only logs when
_moe_layer_loads is populated and a wandb log function is provided.
Parameters:
Current training/benchmark step for wandb x-axis.
Callable like wandb.log or wandb_run.log.
Execute a single training step.
Parameters:
List of batches of training data.
Gradient clipping norm. Optional, if None will not clip gradients.
Run one pass over a single validation dataloader.
Parameters:
Name of the validation dataset.
DataLoader for the validation dataset.
Whether this rank owns the trainable model and its components.
Log metrics to wandb and other loggers.
Parameters:
MetricsSample object, containing: step: int, the current step. epoch: int, the current epoch. metrics: Dict[str, float], containing: “loss”: Training loss. “grad_norm”: Grad norm from the training step. “lr”: Learning rate. “mem”: Memory allocated. “tps”: Tokens per second. “tps_per_gpu”: Tokens per second per GPU. “num_label_tokens”: Number of label tokens.
Log metrics to wandb, MLflow and other loggers Args: log_data: MetricsSample object, containing: step: int, the current step. epoch: int, the current epoch. metrics: Dict[str, float], containing: “val_loss”: Validation loss. “lr”: Learning rate. “num_label_tokens”: Number of label tokens. “mem”: Memory allocated.
Run the training loop over all epochs and batches.
For each batch, perform a forward pass, compute loss, backpropagate, and update model parameters when necessary. Also prints loss every gradient step.
Builds all components needed for training/validation/logging/checkpointing/etc.
This is the last place where self.cfg should be referenced.
Raises:
NotImplemented: Raises if it tries to restore a checkpoint; will be removed.
Build partial CUDA-graph state only for explicitly enabled model backends.
Parameters:
Fully initialized model roots or pipeline-local model parts.
Whether PyTorch activation checkpointing is enabled.
Whether pipeline parallelism is enabled.
Returns: PartialCudaGraphManager | None
An armed manager when any backend selects CUDA-graph scopes, otherwise None.
Return a collate-fn wrapper that precomputes pipeline-parallel causal masks, or None.
None when PP is disabled, the model config can’t be loaded, or the model
computes masks internally (e.g. deepseek_v4 or glm_moe_dsa). Passed to
DataloaderConfig.build as collate_wrapper.
Extract the explicit Megatron training blend for domain-mixture construction.
Downgrade to MaskedCrossEntropy when the requested loss cannot run.
Return whether validation must use the configured training packer.
Return whether the recipe should attach PP causal-mask precomputation.
Return whether loss_fn accepts the per-token loss_weights contract.
Build and initialize a model.
Parameters:
Configuration for model instantiation.
Configuration for PEFT.
Random seed.
Whether using packed sequences.
Configuration for FP8.
Configuration for torch.compile.
Configuration for BitsAndBytes quantization.
Resolved distributed topology and policy object.
Configuration for QAT (will be instantiated to QATConfig).
Freeze configuration (freeze_config YAML section as a
mapping, or a typed FreezeConfig) controlling parameter trainability.
Explicit list of SDPA backend name strings (e.g.
["flash_attention", "efficient_attention"]), or None to
auto-select based on CP / activation checkpointing.
Pre-created device mesh forwarded when distributed_setup is not provided.
Compute the value of trust_remote_code based on the model configuration.
Parameters:
Model configuration.
Returns:
Whether to trust remote code.
Main entry point for the fine-tuning recipe.
Loads the configuration, sets up the trainer, and initiates the training loop.