nemo_automodel.recipes.llm.train_ft

View as Markdown

Module Contents

Classes

NameDescription
TrainFinetuneRecipeForNextTokenPredictionRecipe for fine-tuning a model for next-token prediction.

Functions

NameDescription
_build_partial_cuda_graph_managerBuild partial CUDA-graph state only for explicitly enabled model backends.
_build_pp_collate_wrapperReturn a collate-fn wrapper that precomputes pipeline-parallel causal masks, or None.
_build_tokenizer-
_get_domain_mixture_blendExtract the explicit Megatron training blend for domain-mixture construction.
_get_model_name-
_maybe_downgrade_loss_fnDowngrade to MaskedCrossEntropy when the requested loss cannot run.
_should_pack_validationReturn whether validation must use the configured training packer.
_should_precompute_pp_causal_masksReturn whether the recipe should attach PP causal-mask precomputation.
_supports_loss_weightsReturn whether loss_fn accepts the per-token loss_weights contract.
build_modelBuild and initialize a model.
compute_trust_remote_code_from_modelCompute the value of trust_remote_code based on the model configuration.
mainMain entry point for the fine-tuning recipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction(
cfg
)

Bases: BaseRecipe

Recipe for fine-tuning a model for next-token prediction.

This class orchestrates training, from setup to main training loop.

cfg
magi
= MagiState()
nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._broadcast_from_last_pp_stage(
tensor: torch.Tensor
) -> torch.Tensor

Broadcast a PP last-stage scalar to the other ranks in its pipeline group.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._collect_moe_load_balance()

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.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._configure_packing() -> nemo_automodel.components.models.common.packing.PackingCapabilities

Configure every local model stage and return its NEAT data requirements.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._configure_pipeline_loss_fn()
nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._create_distributed_setup() -> nemo_automodel.components.distributed.config.DistributedSetup

Create the distributed setup used by this recipe rank.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._enable_qat_if_delayed(
step: int
)
nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._forward_backward_step(
idx,
batch,
loss_buffer,
num_label_tokens,
num_batches,
is_train: bool = True
)

Run one local batch and accumulate its loss and optional gradients.

Parameters:

idx

Microbatch index in the accumulation window.

batch

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].

loss_buffer

List receiving the detached scalar loss.

num_label_tokens

Global supervised-token count for loss normalization.

num_batches

Number of microbatches in the accumulation window.

is_train
boolDefaults to True

Whether to backpropagate the combined main and MTP loss.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._log_moe_metrics(
step: int,
wandb_log_fn
) -> None

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:

step
int

Current training/benchmark step for wandb x-axis.

wandb_log_fn

Callable like wandb.log or wandb_run.log.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._run_train_optim_step(
batches,
max_grad_norm: float | None = None
)

Execute a single training step.

Parameters:

batches

List of batches of training data.

max_grad_norm
float | NoneDefaults to None

Gradient clipping norm. Optional, if None will not clip gradients.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._run_validation_epoch(
val_dataloader
)

Run one pass over a single validation dataloader.

Parameters:

val_name

Name of the validation dataset.

val_dataloader

DataLoader for the validation dataset.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._setup_qat(
cfg,
model_parts: list[torch.nn.Module]
)
nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction._should_setup_training_components() -> bool

Whether this rank owns the trainable model and its components.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction.log_train_metrics(
log_data
)

Log metrics to wandb and other loggers.

Parameters:

log_data

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.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction.log_val_metrics(
val_name,
log_data,
metric_logger = None
)

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.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction.run_train_validation_loop()

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.

nemo_automodel.recipes.llm.train_ft.TrainFinetuneRecipeForNextTokenPrediction.setup()

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.
nemo_automodel.recipes.llm.train_ft._build_partial_cuda_graph_manager(
model_parts: list[torch.nn.Module],
activation_checkpointing: bool,
pipeline_parallel: bool
) -> nemo_automodel.components.cuda_graphs.PartialCudaGraphManager | None

Build partial CUDA-graph state only for explicitly enabled model backends.

Parameters:

model_parts
list[nn.Module]

Fully initialized model roots or pipeline-local model parts.

activation_checkpointing
bool

Whether PyTorch activation checkpointing is enabled.

pipeline_parallel
bool

Whether pipeline parallelism is enabled.

Returns: PartialCudaGraphManager | None

An armed manager when any backend selects CUDA-graph scopes, otherwise None.

nemo_automodel.recipes.llm.train_ft._build_pp_collate_wrapper(
cfg_model,
pp_enabled: bool
)

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.

nemo_automodel.recipes.llm.train_ft._build_tokenizer(
cfg_model,
cfg_ds
)
nemo_automodel.recipes.llm.train_ft._get_domain_mixture_blend(
) -> tuple[list[str], list[float]]

Extract the explicit Megatron training blend for domain-mixture construction.

nemo_automodel.recipes.llm.train_ft._get_model_name(
cfg_model
)
nemo_automodel.recipes.llm.train_ft._maybe_downgrade_loss_fn(
loss_fn: torch.nn.Module,
probe_module: torch.nn.Module,
pp_enabled: bool
) -> torch.nn.Module

Downgrade to MaskedCrossEntropy when the requested loss cannot run.

nemo_automodel.recipes.llm.train_ft._should_pack_validation(
model: torch.nn.Module
) -> bool

Return whether validation must use the configured training packer.

nemo_automodel.recipes.llm.train_ft._should_precompute_pp_causal_masks(
model_config: typing.Any
) -> bool

Return whether the recipe should attach PP causal-mask precomputation.

nemo_automodel.recipes.llm.train_ft._supports_loss_weights(
loss_fn: torch.nn.Module
) -> bool

Return whether loss_fn accepts the per-token loss_weights contract.

nemo_automodel.recipes.llm.train_ft.build_model(
cfg_model,
cfg_peft,
seed,
has_packed_sequence = False,
cfg_fp8 = None,
cfg_compile = None,
cfg_quantization = None,
cfg_qat = None,
sdpa_method: list[str] | None = None,
device_mesh = None
) -> tuple[torch.nn.Module | nemo_automodel.components.distributed.pipelining.AutoPipeline, list['Optimizer']]

Build and initialize a model.

Parameters:

cfg_model

Configuration for model instantiation.

cfg_peft

Configuration for PEFT.

seed

Random seed.

has_packed_sequence
Defaults to False

Whether using packed sequences.

cfg_fp8
Defaults to None

Configuration for FP8.

cfg_compile
Defaults to None

Configuration for torch.compile.

cfg_quantization
Defaults to None

Configuration for BitsAndBytes quantization.

distributed_setup
DistributedSetup | NoneDefaults to None

Resolved distributed topology and policy object.

cfg_qat
Defaults to None

Configuration for QAT (will be instantiated to QATConfig).

cfg_freeze
ConfigNode | dict[str, Any] | FreezeConfig | NoneDefaults to None

Freeze configuration (freeze_config YAML section as a mapping, or a typed FreezeConfig) controlling parameter trainability.

sdpa_method
list[str] | NoneDefaults to None

Explicit list of SDPA backend name strings (e.g. ["flash_attention", "efficient_attention"]), or None to auto-select based on CP / activation checkpointing.

device_mesh
Defaults to None

Pre-created device mesh forwarded when distributed_setup is not provided.

nemo_automodel.recipes.llm.train_ft.compute_trust_remote_code_from_model(
cfg_model
)

Compute the value of trust_remote_code based on the model configuration.

Parameters:

cfg_model
ConfigNode

Model configuration.

Returns:

Whether to trust remote code.

nemo_automodel.recipes.llm.train_ft.main(
config_path = None
)

Main entry point for the fine-tuning recipe.

Loads the configuration, sets up the trainer, and initiates the training loop.

nemo_automodel.recipes.llm.train_ft.logger = logging.getLogger(__name__)