nemo_automodel.recipes.llm.train_eagle3
nemo_automodel.recipes.llm.train_eagle3
EAGLE-3 training recipe for Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).
Module Contents
Classes
Functions
Data
API
Bases: PeagleRecipeMixin, BaseRecipe
Recipe for EAGLE-3 training on Llama-style dense LLMs (Llama, Phi-3, Qwen3) and MoE backbones (Qwen3-MoE).
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.
Build the checkpointer using the same plumbing as the standard recipes.
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).
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.
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.
Restore EAGLE-3 meta: global_step, epoch, and vocab mapping tensors.
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_args config node.
Checkpoint path or hub id of the frozen target.
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.
Log a saved checkpoint on rank 0 when checkpointing is enabled.
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.
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.
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).
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:
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.
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.
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.
Validate and cap the prefetch depth for the configured backend.
Persist EAGLE-3 meta: global_step, epoch, and vocab mapping tensors.
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.
Load the target on this GPU and capture supervision in-process.
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.
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.
Connect to one or more remote target servers (no target loaded here).
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.
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.
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.
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, 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.
Run the minimal EAGLE-3 train loop.
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.
Build target model, draft model, data, optimizer, and trainer module.
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).
Run a teardown step, logging (never raising) on failure so one failed step does not abort the rest of cleanup.
Build the backend used by the frozen expert-parallel Kimi K3 target.
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.
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.
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.
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.
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.
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 forcesis_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 interceptF.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 bycp_sizeapplies.
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.
Reject the combinations a Kimi K3 target cannot serve, before it is loaded.
Parameters:
recipe_args.target_model_backend.
distributed.cp_size.
distributed.tp_size.
distributed.pp_size.
recipe_args.packed_sequence_size.
recipe_args.parallel_drafting (P-EAGLE).
recipe_args.target_force_hf.
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.
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.
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.
Main entry point for the EAGLE-3 recipe.