nemo_automodel.recipes.base_recipe
nemo_automodel.recipes.base_recipe
Module Contents
Classes
Functions
Data
API
BaseRecipe provides checkpoint load/save functionality for recipes.
Overriden setattr to keep track of stateful classes.
Parameters:
attribute named.
Value assigned
Raises:
ValueError: if __state_tracked is attemped to be overwriten.
Return the recipe-level autocast context configured by the strategy.
Return the user-facing checkpoint retention policy message, if available.
Return common recipe attributes derived from a distributed setup.
Finalize pending checkpoint publication and always close the checkpointer.
Return globally unique stage indices for local pipeline optimizers.
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).
Load tracked state and return (model, optimizer, scheduler) for downstream loader calls.
Log the checkpoint retention policy without requiring a StepScheduler.
Log metadata and config on main rank using YAML markers.
Log import paths and versions for nemo_automodel, transformers, and torch.
Log model repr, parameter stats, param norm, optimizer and lr scheduler with YAML markers.
Log step scheduler details.
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.
Run manual garbage collection if the current step is configured for it.
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.
Initialize manual garbage collection based on step scheduler config.
Update tqdm bar with loss/lr/tps from a metrics dict (no-op if pbar is 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:
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
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:
The current epoch.
The current step.
The current training loss.
The current validation losses.
The validation metric key used to select the best checkpoint.
Stop tracking one or more attributes for BaseRecipe checkpointing.
Barrier if torch.distributed is initialized. TODO(@akoumpa): deprecate in favor of deviemesh api
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.
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).
True if distributed is not initialized or this process is rank 0. TODO(@akoumpa): deprecate in favor of deviemesh api
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.
Compare two model signatures with normalization so YAML round-trip does not cause false mismatches.
Checks whether object has load_state_dict and state_dict functions.
TODO: also need to check function signatures.
Parameters:
the object to check.
Returns:
returns True if has callable load_state_dict and state_dict
Checks whether object is a dataloader.
Parameters:
the object to check.
Returns:
returns True if object is a dataloader.
Checks whether object should be saved through distributed checkpointing.
Checks whether object is a learning rate scheduler.
Parameters:
the object to check.
Returns:
returns True if object is an OptimizerParamScheduler.
Checks whether object is a model.
Checks whether object is an optimizer.
Checks whether object is a tokenizer or VLM processor.
Parameters:
the object to check.
Returns:
returns True if object is a VLM processor or tokenizer.