nemo_automodel.recipes.llm.train_eagle1
nemo_automodel.recipes.llm.train_eagle1
EAGLE-1 / EAGLE-2 training recipe for Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).
Module Contents
Classes
Functions
Data
API
Bases: BaseRecipe
Recipe for EAGLE-1 training on Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).
Build the checkpointer using the same plumbing as the standard recipes.
Run the frozen target and the draft over one micro-batch.
Parameters:
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
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:
The recipe_args config node.
Target model id or path, used for checkpoint metadata and the default W&B run name.
Prefix for the default W&B run name.
Restore EAGLE-recipe-specific scalars. Subclasses extend this.
Log a saved checkpoint on rank 0 when checkpointing is enabled.
Return the per-term losses this recipe logs alongside the total.
Parameters:
The step metrics returned by _compute_metrics.
Returns: dict[str, float]
Mapping of log-suffix to scalar value, logged as train/<key>.
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.
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).
Persist EAGLE-recipe-specific scalars. Subclasses extend this.
Log a metrics dict to W&B when a run is active (rank 0).
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.
Run the training loop.
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.
Build target model, draft model, data, optimizer, and trainer module.
Build the EAGLE-1/2 augmentation from the recipe’s fixed-width knob.
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.
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).
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 > 1with 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.
Entrypoint for TrainEagle1Recipe.