nemo_automodel.recipes.llm.train_eagle3

View as Markdown

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

Module Contents

Classes

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

Functions

NameDescription
_all_reduce_mean-
_all_reduce_sum-
_apply_draft_peft_and_fp8Apply the optional peft: (LoRA) and fp8: blocks to the freshly built draft.
_best_effortRun a teardown step, logging (never raising) on failure so one failed step
_build_kimi_k3_target_backendBuild the backend used by the frozen expert-parallel Kimi K3 target.
_export_merged_lora_draftWrite the serve-ready consolidated export of a LoRA run (adapters merged into the base draft).
_load_draft_weightsWarm-start the draft from a previously trained draft’s consolidated safetensors export.
_merged_lora_state_dictReturn the full draft state dict with LoRA deltas folded into the base weights.
_submesh_or_noneReturn the named 1D submesh (e.g. “cp”/“dp”) or None if absent.
_validate_cached_eagle3_manifestValidate that an EAGLE-3 offline cache matches the configured target and recipe options.
_validate_cp_gatesReject context-parallel combinations the EAGLE-3 path cannot honor.
_validate_kimi_k3_gatesReject the combinations a Kimi K3 target cannot serve, before it is loaded.
_validate_peagle_gatesReject P-EAGLE (parallel_drafting) combinations its trainer cannot honor.
_validate_tp_gatesReject tensor-parallel combinations the EAGLE-3 target path cannot honor.
_window_tau_simSimulated accept length over a metrics window, reduced across ranks.
mainMain entry point for the EAGLE-3 recipe.

Data

KIMI_K3_MODEL_TYPE

logger

API

class nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe(
cfg
)

Bases: PeagleRecipeMixin, BaseRecipe

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

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._all_reduce_draft_grads_over_cp() -> None

Sum the draft gradients over the cp group before the optimizer step.

The draft runs sequence-sharded across cp, so each rank holds only its shard’s gradient contribution; summing yields the full-sequence gradient (the loss is already globally normalized by the trainer). DDP has averaged over dp, and sum-over-cp / avg-over-dp commute, so the cp replicas end up identical.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._build_checkpointer(
target_path: str
) -> None

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

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._build_train_dataloader(
data_path,
split
)

Build the train dataloader from self.cfg.recipe_args.

Reused by the on-policy regen loop to rebuild the train dataloader against a fresh shard directory with identical settings (only data_path/split vary).

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._finalize_training() -> None

Release training resources on any exit path (normal, early-return, or exception). Best-effort: each step is guarded so a failure in one does not block the others.

The high-value step is disconnecting the remote target. Without it a mid-training crash leaves the long-lived target server with a stale client-idle state and a half-open NCCL transport, so the next run cannot connect. close() is a no-op for the co-located backend. The process group is intentionally left alone — it is a framework-global resource that direct callers (tests, the interactive launcher) reuse after the loop returns, and initialize_distributed already destroys it at process exit.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._forward_batch(
batch,
target_batch = None
)

Run the trainer module for one batch, from the live target or the cache.

target_batch may be supplied when the supervision was prefetched asynchronously (remote backend); it is already on the training device and self-contained, so the raw batch is not needed.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._load_extra_state(
ckpt_dir: str
) -> None

Restore EAGLE-3 meta: global_step, epoch, and vocab mapping tensors.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._load_kimi_k3_target(
recipe_cfg,
target_path,
text_config
)

Load Kimi K3 as a frozen expert-parallel target.

Kimi K3 has no HuggingFace implementation and its 896 routed experts do not fit on one GPU, so the target is built from the (possibly nested) text config through AutoModel’s custom path with an explicit expert-parallel backend instead of the generic from_pretrained call. The EAGLE-3 draft only ever consumes the text decoder’s hidden states, so a vision-language checkpoint is loaded as its text-only KimiK3ForCausalLM.

Parameters:

recipe_cfg

recipe_args config node.

target_path

Checkpoint path or hub id of the frozen target.

text_config

The target’s text config (config.text_config when the checkpoint is vision-language, otherwise the config itself).

Returns:

The frozen Kimi K3 target model.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._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_eagle3.TrainEagle3Recipe._maybe_run_decode_eval(
launch: bool = True
) -> None

Rank-0 log-point hook: collect finished decode-eval results, launch the next one when due.

Runs outside any collective path (pure subprocess/file I/O), so only rank 0 calls it. Results are logged at the CURRENT optimizer step (wandb steps must not go backwards); train/tau_real_step records the step the evaluated snapshot was actually taken at.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._maybe_run_regen() -> None

Window-boundary hook: launch a regeneration cycle when the cadence is due.

Called at every grad-accum window boundary (not just log points) so the launch cadence follows regen.every_steps regardless of log_every_steps. A no-op on non-main ranks, where regen_runner is None.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._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_eagle3.TrainEagle3Recipe._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_eagle3.TrainEagle3Recipe._maybe_shard_cp(
inputs: dict
) -> dict

Shard the draft supervision along the sequence for context parallelism.

The target emits full-sequence aux/logits (gathered); the draft runs sharded, so every [batch, sequence, ...] tensor in inputs is index-selected to this rank’s cp shard (via a global index). Non-packed runs then get global position_ids (the shard’s indices); packed runs keep their sliced per-document position_ids. The index is the contiguous shard by default, or the load-balanced zig-zag layout (rank r owns chunks r and 2*cp-1-r) when cp_zigzag. No-op without CP.

Parameters:

inputs
dict

Trainer-input mapping; tensor values shaped [batch, sequence, ...] (matching the full sequence length) are sharded along the sequence axis, others pass through unchanged.

Returns: dict

The same mapping with sequence-length tensors replaced by their cp shard.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._maybe_swap_regen_dataloader() -> bool

Swap the train dataloader to the newest ready regenerated shard dir; return whether a swap happened.

Rank 0 decides which directory (if any) is ready and broadcasts it; every rank then rebuilds its dataloader against the same shared-filesystem path so the distributed sampler stays consistent. Called at a grad-accum window boundary (pending_micro_batches == 0), so there is no partial window to flush; the caller ends the current segment and rebinds when this returns True. The regenerated slice is expected to keep the original prompt count so steps-per-epoch (and the fixed LR horizon) stay coherent; a differing length is warned about rather than silently trusted.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._module()
nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._prefetched_batches(
dataloader
)

Yield (batch, target_batch) keeping up to target_prefetch_depth remote target requests in flight, so target inference on the server(s) overlaps draft training on this GPU.

Requests are dispatched round-robin across servers by the backend; the depth is capped to the server count (see setup) so each server has at most one in-flight request — required for NCCL recv ordering.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._resolve_prefetch_depth(
recipe_cfg
) -> int

Validate and cap the prefetch depth for the configured backend.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._run_eval()
nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._save_extra_state(
path: str,
epoch: int
) -> None

Persist EAGLE-3 meta: global_step, epoch, and vocab mapping tensors.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_cached_target(
recipe_cfg,
draft_base_config
)

Offline path: stream a precomputed cache; no target model is loaded.

Reads the cache manifest for the draft-vocab mapping and loads the stored target embeddings (the one target tensor the draft still needs). Sets self.target_model / self.target_wrapper to None and builds the cache-backed dataloader.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_colocated_target(
recipe_cfg,
target_path,
draft_base_config
)

Load the target on this GPU and capture supervision in-process.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_online_target(
recipe_cfg,
target_path,
draft_base_config
)

Live path: load the target model and build the live dataloader.

Sets self.target_model / self.target_wrapper / self.train_dataloader / self.val_dataloader and returns the (selected_token_ids, selected_token_mask) draft-vocab mapping built by scanning the training data.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_regen(
recipe_cfg,
target_path
) -> None

Build the optional on-policy regeneration runner (rank 0, online path only).

The runner is only constructed on rank 0 (it owns the worker subprocess); every rank participates in the epoch-boundary dataloader swap via the broadcast in _maybe_swap_regen_dataloader, so all ranks must know whether the feature is enabled. _regen_enabled carries that flag.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_remote_target(
recipe_cfg
)

Connect to one or more remote target servers (no target loaded here).

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_sglang_target(
recipe_cfg,
target_path
)

Co-located SGLang target: serve the frozen target through SGLang on this GPU.

Same supervision contract as colocated (full-vocab logits shipped to the trainer, draft-vocab projection trainer-side), but the target forward runs through SGLang’s ModelRunner. SGLang carves its weight + KV pool out of this GPU up front (mem_fraction_static), and the draft trains in the remainder.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._setup_vllm_target(
recipe_cfg,
target_path
)

Co-located vLLM target: serve the frozen target through vLLM on this GPU.

Same supervision contract as colocated (full-vocab logits shipped to the trainer, draft-vocab projection trainer-side), but the target forward runs through vLLM’s extract_hidden_states path. vLLM carves its weight

  • KV pool out of this GPU up front (gpu_memory_utilization), and the draft trains in the remainder.
nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._train_epochs(
start_epoch,
is_ddp,
pbar = None
)

Run the epoch loop (extracted so run_train_validation_loop can wrap it in try/finally and guarantee teardown on any exit path).

The run is step-budget driven: it stops when global_step reaches total_optim_steps (the horizon the LR schedule was calibrated on), not when the epoch range is exhausted. With on-policy regen enabled an epoch is split into segments — one contiguous pass over each bound dataloader — and a swap to freshly regenerated shards ends the current segment and rebinds a new one. Swaps happen only at a grad-accum window boundary (pending_micro_batches == 0), so there is never a partial window to flush or an un-synced gradient at the swap point; the trailing flush below then only ever runs for the single segment that ends by dataloader exhaustion.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe._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_eagle3.TrainEagle3Recipe.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, global_step, and the EAGLE-3 vocab mapping tensors. 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_eagle3.TrainEagle3Recipe.run_train_validation_loop()

Run the minimal EAGLE-3 train loop.

nemo_automodel.recipes.llm.train_eagle3.TrainEagle3Recipe.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-3 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. The EAGLE-3 vocab mapping tensors ride along through _save_extra_state.

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_eagle3.TrainEagle3Recipe.setup()

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

nemo_automodel.recipes.llm.train_eagle3._all_reduce_mean(
value: torch.Tensor
) -> torch.Tensor
nemo_automodel.recipes.llm.train_eagle3._all_reduce_sum(
value: torch.Tensor
) -> torch.Tensor
nemo_automodel.recipes.llm.train_eagle3._apply_draft_peft_and_fp8(
draft_model,
cfg,
parallel_drafting: bool,
freeze_embeddings: bool = True
)

Apply the optional peft: (LoRA) and fp8: blocks to the freshly built draft.

LoRA freezes every base draft parameter and adds trainable lora_A/lora_B adapters to the matched linears; meaningful only when the draft was warm-started from a trained checkpoint (recipe_args.draft_weights_path). P-EAGLE is rejected because parallel drafting must train mask_hidden and the embeddings, which the LoRA freeze would lock. FP8 + LoRA is rejected like QAT + PEFT in the SFT recipe. Modifies the draft in place and returns the instantiated PeftConfig (or None).

nemo_automodel.recipes.llm.train_eagle3._best_effort(
label: str,
fn
) -> None

Run a teardown step, logging (never raising) on failure so one failed step does not abort the rest of cleanup.

nemo_automodel.recipes.llm.train_eagle3._build_kimi_k3_target_backend(
recipe_cfg

Build the backend used by the frozen expert-parallel Kimi K3 target.

nemo_automodel.recipes.llm.train_eagle3._export_merged_lora_draft(
draft_model,
path: str
) -> str

Write the serve-ready consolidated export of a LoRA run (adapters merged into the base draft).

Restores the invariant full-FT runs have (save_consolidated): the final checkpoint of a LoRA run contains model/consolidated with full merged weights (plus the d2t/t2d buffers riding along in the state dict) and the draft config.json, so serve/bench and a later draft_weights_path warm start can consume it without any external merge step.

nemo_automodel.recipes.llm.train_eagle3._load_draft_weights(
draft_model: torch.nn.Module,
path: str
) -> None

Warm-start the draft from a previously trained draft’s consolidated safetensors export.

path is either a single .safetensors file or a directory (flat or model.safetensors.index.json-sharded, the layouts save_consolidated writes). Loaded with strict=False because the vocab-mapping buffers (d2t/t2d) are regenerated by set_vocab_mapping each run and need not be present.

nemo_automodel.recipes.llm.train_eagle3._merged_lora_state_dict(
draft_model
) -> dict[str, torch.Tensor]

Return the full draft state dict with LoRA deltas folded into the base weights.

For every LoRA-patched linear, W += (alpha / dim) * lora_B.weight @ lora_A.weight (the eval-time equivalent of the adapter pathway). Under DoRA the effective weight is additionally row-rescaled to the learned magnitude, W' = (lora_magnitude / ||W + scale * B @ A||_row) * (W + scale * B @ A) (matching LinearLoRA’s DoRA forward at eval time). lora_* keys (including lora_magnitude) are dropped.

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

Return the named 1D submesh (e.g. “cp”/“dp”) or None if absent.

Uses get_flat_mesh so _flatten()-created axes (“dp”) resolve on PyTorch 2.9-2.11, where a plain device_mesh[name] is deprecated or the name is missing from mesh_dim_names.

nemo_automodel.recipes.llm.train_eagle3._validate_cached_eagle3_manifest(
cache_dir: str,
manifest: dict,
draft_base_config,
mask_reasoning_content: bool,
mask_generation_prompt: bool
) -> None

Validate that an EAGLE-3 offline cache matches the configured target and recipe options.

The cached trainer streams the precomputed aux features, draft-vocab targets and loss_mask/position_mask as stored, so the target’s vocabulary and width and the recipe’s mask options must be the ones the producer (precompute_eagle3) ran with; a mismatch would otherwise crash deep inside the draft or, for the mask, train silently on the wrong tokens.

nemo_automodel.recipes.llm.train_eagle3._validate_cp_gates(
cp_size: int,
backend: str,
packed_sequence_size: int,
target_force_hf: bool = False,
target_attn_implementation: str | None = None,
seq_length: int | None = None,
cp_zigzag: bool = False,
cp_mode: str = 'ring'
) -> None

Reject context-parallel combinations the EAGLE-3 path cannot honor.

Two CP compute modes are supported for the draft:

  • "ring" shards the frozen target too and forces is_causal (the self_attn K/V-gather hook strips the attention_mask), so it is incompatible with sequence packing (which needs the 4D block-causal mask); it pins the flash-attn 2.8.x ring kernels and needs the HF-SDPA target so the hook can intercept F.scaled_dot_product_attention.
  • "ulysses" runs the draft attention as an all-to-all over the full (gathered) sequence, so it composes with packing and leaves the target on its normal forward — only sequence divisibility by cp_size applies.

The remote backend is unsupported under CP in either mode (the target runs out-of-process). Preconditions are checked here so a misconfig fails at setup rather than mid-forward.

nemo_automodel.recipes.llm.train_eagle3._validate_kimi_k3_gates(
backend: str,
cp_size: int,
tp_size: int,
pp_size: int,
packed_sequence_size: int,
parallel_drafting: bool,
target_force_hf: bool
) -> None

Reject the combinations a Kimi K3 target cannot serve, before it is loaded.

Parameters:

backend
str

recipe_args.target_model_backend.

cp_size
int

distributed.cp_size.

tp_size
int

distributed.tp_size.

pp_size
int

distributed.pp_size.

packed_sequence_size
int

recipe_args.packed_sequence_size.

parallel_drafting
bool

recipe_args.parallel_drafting (P-EAGLE).

target_force_hf
bool

recipe_args.target_force_hf.

nemo_automodel.recipes.llm.train_eagle3._validate_peagle_gates(
backend: str,
cached_target_path,
packed_sequence_size: int,
lk_loss_type = None
) -> None

Reject P-EAGLE (parallel_drafting) combinations its trainer cannot honor.

PEagleTrainerModule.forward consumes the live colocated target’s full-vocab target_logits only. The remote and offline-cache backends instead supply precomputed draft-vocab target_probs/position_mask (a parameter mismatch), and sequence packing feeds position_ids/seq_lens/ doc_remaining the P-EAGLE forward does not accept (and the partitioned path would run the target on a packed row without per-document masking, leaking across documents). P-EAGLE only safely supports a colocated live target on non-packed sequences.

nemo_automodel.recipes.llm.train_eagle3._validate_tp_gates(
tp_size: int,
backend: str,
cp_size: int
) -> None

Reject tensor-parallel combinations the EAGLE-3 target path cannot honor.

TP shards only the colocated target (the FSDP2 parallelize plan column/row shards its linears and makes the lm_head logits a vocab-sharded DTensor, which the target wrapper gathers). It is therefore meaningless with the remote backend (the target runs out-of-process), and combined TP+CP is not yet wired: the CP gather path does not handle a TP-sharded (DTensor) sequence.

nemo_automodel.recipes.llm.train_eagle3._window_tau_sim(
step_prefix_hits: torch.Tensor | None,
step_valid: torch.Tensor | None
) -> float | None

Simulated accept length over a metrics window, reduced across ranks.

step_prefix_hits / step_valid are the window’s accumulated per-TTT-step prefix-hit / supervised counts (None when the trainer does not report them, e.g. P-EAGLE). Counts are extensive, so they are sum-reduced across ranks before forming the per-step survival rates; every rank must therefore call this at the same point. Returns None when there is nothing to report.

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

Main entry point for the EAGLE-3 recipe.

nemo_automodel.recipes.llm.train_eagle3.KIMI_K3_MODEL_TYPE = KimiK3TextConfig.model_type
nemo_automodel.recipes.llm.train_eagle3.logger = logging.getLogger(__name__)