nemo_automodel.recipes.base_recipe

View as Markdown

Module Contents

Classes

NameDescription
BaseRecipeBaseRecipe provides checkpoint load/save functionality for recipes.

Functions

NameDescription
_dist_barrierBarrier if torch.distributed is initialized.
_extract_model_signatureExtract a stable subset of the model config used to decide checkpoint compatibility.
_is_checkpoint_model_config_compatibleCompare the checkpoint’s saved config.yaml model signature to the
_is_rank_0True if distributed is not initialized or this process is rank 0.
_normalize_signature_valueNormalize a signature value so that YAML round-trip and minor type differences
_signatures_matchCompare two model signatures with normalization so YAML round-trip does not cause false mismatches.
has_load_restore_stateChecks whether object has load_state_dict and state_dict functions.
is_dataloaderChecks whether object is a dataloader.
is_distributed_statefulChecks whether object should be saved through distributed checkpointing.
is_lr_schedulerChecks whether object is a learning rate scheduler.
is_modelChecks whether object is a model.
is_optimizerChecks whether object is an optimizer.
is_tokenizerChecks whether object is a tokenizer or VLM processor.

Data

logger

API

class nemo_automodel.recipes.base_recipe.BaseRecipe()

BaseRecipe provides checkpoint load/save functionality for recipes.

nemo_automodel.recipes.base_recipe.BaseRecipe.__setattr__(
key,
value
)

Overriden setattr to keep track of stateful classes.

Parameters:

key
str

attribute named.

value
Any

Value assigned

Raises:

  • ValueError: if __state_tracked is attemped to be overwriten.
nemo_automodel.recipes.base_recipe.BaseRecipe._autocast_context()

Return the recipe-level autocast context configured by the strategy.

nemo_automodel.recipes.base_recipe.BaseRecipe._checkpoint_retention_policy_message(
checkpoint_config = None
) -> str | None

Return the user-facing checkpoint retention policy message, if available.

nemo_automodel.recipes.base_recipe.BaseRecipe._distributed_setup_attributes(
distributed_setup
)
staticmethod

Return common recipe attributes derived from a distributed setup.

nemo_automodel.recipes.base_recipe.BaseRecipe._dp_allreduce(
tensor,
op = dist.ReduceOp.SUM,
include_cp: bool = False
)
nemo_automodel.recipes.base_recipe.BaseRecipe._finalize_and_close_checkpointer() -> None

Finalize pending checkpoint publication and always close the checkpointer.

nemo_automodel.recipes.base_recipe.BaseRecipe._get_cp_group_size()
nemo_automodel.recipes.base_recipe.BaseRecipe._get_dp_group(
include_cp: bool = False
)
nemo_automodel.recipes.base_recipe.BaseRecipe._get_dp_group_size(
include_cp: bool = False
)
nemo_automodel.recipes.base_recipe.BaseRecipe._get_dp_rank(
include_cp: bool = False
)
nemo_automodel.recipes.base_recipe.BaseRecipe._get_optimizer_checkpoint_part_ids() -> list[int] | None

Return globally unique stage indices for local pipeline optimizers.

nemo_automodel.recipes.base_recipe.BaseRecipe._get_pp_group()

Return the pipeline-parallel process group, or None when pp is disabled.

Threaded to the checkpointer so PEFT adapters are gathered across PP stages at save time; without it the on-disk adapter only contains the local stage’s layers (see _gather_peft_state_dict_across_pp).

nemo_automodel.recipes.base_recipe.BaseRecipe._get_pp_rank()
nemo_automodel.recipes.base_recipe.BaseRecipe._get_tp_rank()
nemo_automodel.recipes.base_recipe.BaseRecipe._load_checkpoint_tracked_state(
ckpt_dir: str
)

Load tracked state and return (model, optimizer, scheduler) for downstream loader calls.

nemo_automodel.recipes.base_recipe.BaseRecipe._log_checkpoint_retention_policy(
checkpoint_config = None
) -> None

Log the checkpoint retention policy without requiring a StepScheduler.

nemo_automodel.recipes.base_recipe.BaseRecipe._log_experiment_details()

Log metadata and config on main rank using YAML markers.

nemo_automodel.recipes.base_recipe.BaseRecipe._log_library_versions()

Log import paths and versions for nemo_automodel, transformers, and torch.

nemo_automodel.recipes.base_recipe.BaseRecipe._log_model_and_optimizer_details(
model: torch.nn.Module | list[torch.nn.Module] | None = None,
optimizer: torch.optim.Optimizer | list[torch.optim.Optimizer] | None = None,
)

Log model repr, parameter stats, param norm, optimizer and lr scheduler with YAML markers.

nemo_automodel.recipes.base_recipe.BaseRecipe._log_step_scheduler_details(
)

Log step scheduler details.

nemo_automodel.recipes.base_recipe.BaseRecipe._make_progress_bar(
total: int | None = None,
initial: int = 0
)

Create a tqdm progress bar on rank 0; returns None on other ranks.

Without arguments the totals come from self.step_scheduler; recipes without a step scheduler (e.g. the EAGLE family) pass total and initial explicitly.

nemo_automodel.recipes.base_recipe.BaseRecipe._maybe_collect_garbage() -> None

Run manual garbage collection if the current step is configured for it.

nemo_automodel.recipes.base_recipe.BaseRecipe._set_moe_aux_loss_backward_scale(
num_batches: int,
num_label_tokens: int
) -> None

Set the per-microbatch MoE auxiliary-loss scale for one optimizer step.

The base scale averages accumulation microbatches and restores the CP sum lost in the flattened DP-CP gradient average. PP additionally needs to compensate for its post-backward token normalization.

nemo_automodel.recipes.base_recipe.BaseRecipe._setup_garbage_collection(
) -> None

Initialize manual garbage collection based on step scheduler config.

nemo_automodel.recipes.base_recipe.BaseRecipe._update_progress_bar(
pbar,
metrics: dict
) -> None

Update tqdm bar with loss/lr/tps from a metrics dict (no-op if pbar is None).

nemo_automodel.recipes.base_recipe.BaseRecipe.load_checkpoint(
restore_from: str | None = None
)

Loads checkpoint with automatic compatibility checking.

This method will:

  • If restore_from is set to a path or “LATEST”: resolve and load that checkpoint
  • If restore_from is None: auto-detect the latest checkpoint in checkpoint_dir
  • Before loading, check if the checkpoint is compatible with the current model config
  • If incompatible: print a warning and proceed with the restore anyway

Parameters:

restore_from
str | NoneDefaults to None

Path to checkpoint directory to restore from. Options:

  • None: Auto-detect latest checkpoint in checkpoint_dir
  • “LATEST”: Explicitly auto-detect latest checkpoint
  • “epoch_0_step_100”: Subdirectory name (relative to checkpoint_dir)
  • ”./path/to/checkpoint”: Absolute or relative path
nemo_automodel.recipes.base_recipe.BaseRecipe.save_checkpoint(
epoch: int,
step: int,
train_loss: float,
val_loss: dict[str, float] | None = None,
best_metric_key: str = 'default'
)

Save the current training state as a checkpoint.

As long as the object has a ‘load_state_dict’ and ‘state_dict’ function, it will be saved.

Parameters:

epoch
int

The current epoch.

step
int

The current step.

train_loss
float

The current training loss.

val_loss
dict[str, float]Defaults to None

The current validation losses.

best_metric_key
strDefaults to 'default'

The validation metric key used to select the best checkpoint.

nemo_automodel.recipes.base_recipe.BaseRecipe.untrack_state(
keys: str = ()
) -> None

Stop tracking one or more attributes for BaseRecipe checkpointing.

nemo_automodel.recipes.base_recipe._dist_barrier(
group = None
) -> None

Barrier if torch.distributed is initialized. TODO(@akoumpa): deprecate in favor of deviemesh api

nemo_automodel.recipes.base_recipe._extract_model_signature(
cfg: dict
) -> dict

Extract a stable subset of the model config used to decide checkpoint compatibility.

This includes model architecture fields AND training-mode indicators (e.g. PEFT) that affect the checkpoint format.

nemo_automodel.recipes.base_recipe._is_checkpoint_model_config_compatible(
current_cfg,
ckpt_dir: str
) -> tuple[bool, str]

Compare the checkpoint’s saved config.yaml model signature to the current run’s model signature.

Uses the effective YAML-ready config (when available) for comparison so runtime overrides are checked against the same representation saved in the checkpoint. Round-tripping through YAML preserves types, avoiding false mismatches that would arise from using to_dict() (which may apply type conversions).

nemo_automodel.recipes.base_recipe._is_rank_0() -> bool

True if distributed is not initialized or this process is rank 0. TODO(@akoumpa): deprecate in favor of deviemesh api

nemo_automodel.recipes.base_recipe._normalize_signature_value(
v
)

Normalize a signature value so that YAML round-trip and minor type differences (e.g. int vs str, ConfigNode vs dict) do not cause false mismatches.

nemo_automodel.recipes.base_recipe._signatures_match(
cur_sig: dict,
ckpt_sig: dict
) -> bool

Compare two model signatures with normalization so YAML round-trip does not cause false mismatches.

nemo_automodel.recipes.base_recipe.has_load_restore_state(
object
)

Checks whether object has load_state_dict and state_dict functions.

TODO: also need to check function signatures.

Parameters:

object
any

the object to check.

Returns:

returns True if has callable load_state_dict and state_dict

nemo_automodel.recipes.base_recipe.is_dataloader(
object
)

Checks whether object is a dataloader.

Parameters:

object
any

the object to check.

Returns:

returns True if object is a dataloader.

nemo_automodel.recipes.base_recipe.is_distributed_stateful(
object
)

Checks whether object should be saved through distributed checkpointing.

nemo_automodel.recipes.base_recipe.is_lr_scheduler(
object
)

Checks whether object is a learning rate scheduler.

Parameters:

object
any

the object to check.

Returns:

returns True if object is an OptimizerParamScheduler.

nemo_automodel.recipes.base_recipe.is_model(
object
)

Checks whether object is a model.

nemo_automodel.recipes.base_recipe.is_optimizer(
object
)

Checks whether object is an optimizer.

nemo_automodel.recipes.base_recipe.is_tokenizer(
object
)

Checks whether object is a tokenizer or VLM processor.

Parameters:

object
any

the object to check.

Returns:

returns True if object is a VLM processor or tokenizer.

nemo_automodel.recipes.base_recipe.logger = logging.getLogger(__name__)