nemo_automodel.recipes.llm.train_eagle1

View as Markdown

EAGLE-1 / EAGLE-2 training recipe for Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).

Module Contents

Classes

NameDescription
TrainEagle1RecipeRecipe for EAGLE-1 training on Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).

Functions

NameDescription
_all_reduce_mean-
_build_feature_noiseBuild the EAGLE-1/2 augmentation from the recipe’s fixed-width knob.
_packing_kwargsSequence-packing metadata from a dataloader batch (empty dict when unpacked).
_submesh_or_noneReturn the named (flattened) submesh, or None if absent / no mesh.
_validate_packing_gatesReject sequence-packing configs the EAGLE-1/2 path cannot honor (fail fast at setup).
mainEntrypoint for TrainEagle1Recipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe(
cfg
)

Bases: BaseRecipe

Recipe for EAGLE-1 training on Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._build_checkpointer(
target_path: str
) -> None

Build the checkpointer using the same plumbing as the standard recipes.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._compute_metrics(
batch: dict[str, torch.Tensor]
)

Run the frozen target and the draft over one micro-batch.

Parameters:

batch
dict[str, torch.Tensor]

Dataloader batch whose tensors are already on self.device. Carries input_ids / attention_mask / loss_mask, each of shape [batch, sequence], plus any packing metadata.

Returns:

EagleStepMetrics for this micro-batch. Subclasses that feed the draft

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._finalize_setup(
recipe_cfg,
target_path: str,
wandb_name_prefix: str
) -> None

Build the optimizer, schedule, checkpointer, and logging around a ready trainer module.

Runs once self.trainer_module / self.draft_model / the dataloaders exist, so recipes that assemble those differently (e.g. ViSpec’s VLM target) share the rest of setup() instead of copying it.

Parameters:

recipe_cfg

The recipe_args config node.

target_path
str

Target model id or path, used for checkpoint metadata and the default W&B run name.

wandb_name_prefix
str

Prefix for the default W&B run name.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._load_extra_state(
ckpt_dir: str
) -> None

Restore EAGLE-recipe-specific scalars. Subclasses extend this.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._log_saved_checkpoint(
kind: str,
epoch: int,
step: int
) -> None

Log a saved checkpoint on rank 0 when checkpointing is enabled.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._loss_components(
metrics
) -> dict[str, float]

Return the per-term losses this recipe logs alongside the total.

Parameters:

metrics

The step metrics returned by _compute_metrics.

Returns: dict[str, float]

Mapping of log-suffix to scalar value, logged as train/<key>.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._maybe_save_final_checkpoint(
completed_epochs: int
) -> bool

Always save the fully-trained model at the end of a completed run, unless a periodic checkpoint already captured the final step.

The end-of-run state is otherwise easy to lose: with no cadence nothing is saved at all, and with a pure step cadence the final step is skipped whenever the total step count is not a multiple of ckpt_every_steps. This is a no-op only when a step or epoch checkpoint already landed on the final step, so it never duplicates or collides with one.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._maybe_save_step_checkpoint(
epoch: int
) -> bool

Save a checkpoint mid-epoch when ckpt_every_steps is configured.

Called after every optimizer step. Saves whenever ckpt_every_steps is a positive integer and the current global_step is a multiple of it. Returns True if a checkpoint was written. The checkpoint directory is named epoch_{epoch}_step_{global_step} so it never collides with the end-of-epoch checkpoint (which uses epoch + 1).

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._module()
nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._run_eval()
nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._save_extra_state(
path: str,
epoch: int
) -> None

Persist EAGLE-recipe-specific scalars. Subclasses extend this.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe._wandb_log(
data: dict,
step: int
) -> None

Log a metrics dict to W&B when a run is active (rank 0).

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe.load_checkpoint(
restore_from: str | None = None
) -> None

Resolve and restore a checkpoint produced by save_checkpoint.

Restores the draft model, optimizer, LR scheduler, RNG, and global_step. Target model weights are NOT restored — they are re-loaded from the HF hub on each run because the target is frozen.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe.run_train_validation_loop()

Run the training loop.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe.save_checkpoint(
epoch: int,
step: int,
train_loss: float | None = None,
val_loss: dict[str, float] | None = None,
best_metric_key: str = 'default',
is_final_checkpoint: bool = False
) -> None

Persist draft model, optimizer, scheduler, RNG, and EAGLE meta.

Overrides BaseRecipe.save_checkpoint because EAGLE recipes hold multiple nn.Module attributes (frozen target, target wrapper, trainer module wrapping the draft) — only draft_model should be persisted as the main model.

is_final_checkpoint is computed by the caller (this hand-rolled loop has no step_scheduler for the checkpointer to infer it from); save_consolidated: final exports HF safetensors only when it is True.

nemo_automodel.recipes.llm.train_eagle1.TrainEagle1Recipe.setup()

Build target model, draft model, data, optimizer, and trainer module.

nemo_automodel.recipes.llm.train_eagle1._all_reduce_mean(
value: torch.Tensor
) -> torch.Tensor
nemo_automodel.recipes.llm.train_eagle1._build_feature_noise(
half_width: float

Build the EAGLE-1/2 augmentation from the recipe’s fixed-width knob.

nemo_automodel.recipes.llm.train_eagle1._packing_kwargs(
batch: dict[str, torch.Tensor]
) -> dict[str, torch.Tensor]

Sequence-packing metadata from a dataloader batch (empty dict when unpacked).

The packed loader (packed_sequence_size > 0) emits position_ids / seq_lens / doc_remaining alongside input_ids; the default loader does not. Keyed on seq_lens so the caller can splat the result into HFEagleTargetModel.generate_batch unconditionally.

nemo_automodel.recipes.llm.train_eagle1._submesh_or_none(
device_mesh,
name: str
)

Return the named (flattened) submesh, or None if absent / no mesh.

Uses get_flat_mesh so _flatten()-created axes (“dp”) resolve across torch versions. The “dp” axis excludes “tp”, so keying the draft DDP group and the dataloader sampler on it replicates the draft across tensor-parallel ranks (every TP rank in a draft replica sees the same batch).

nemo_automodel.recipes.llm.train_eagle1._validate_packing_gates(
cp_size: int,
target_attn_impl: str,
micro_batch_size: int
) -> None

Reject sequence-packing configs the EAGLE-1/2 path cannot honor (fail fast at setup).

  • Context parallelism shards the sequence and strips the 4D block-causal mask packing relies on, and EAGLE-1/2 has no CP sequence-sharding path, so cp_size > 1 with packing would silently train on wrong supervision.
  • A FlashAttention target infers document boundaries from per-document position_ids, which transformers packs only at batch size 1.
nemo_automodel.recipes.llm.train_eagle1.main(
config_path: str | None = None
)

Entrypoint for TrainEagle1Recipe.

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