nemo_automodel.recipes.llm.peagle_recipe
nemo_automodel.recipes.llm.peagle_recipe
P-EAGLE recipe-level logic, split out of the EAGLE-3 recipe.
PeagleRecipeMixin holds the P-EAGLE-only step methods so TrainEagle3Recipe
keeps only the shared EAGLE-3 training flow plus the parallel_drafting /
_peagle_partitioned dispatch. The mixin relies on the recipe attributes it is
mixed into (device, target_wrapper, trainer_module,
grad_accumulation_steps, _module).
Module Contents
Classes
API
P-EAGLE setup and sequence-partitioning step methods for TrainEagle3Recipe.
Validate P-EAGLE recipe args and populate draft_config; return mask_token_id.
Mutates draft_config in place with the P-EAGLE keys (mask_token_id,
num_depths, COD ratios, draft num_hidden_layers) that are serialized
into the saved draft config.json so the checkpoint loads into vLLM’s
parallel-drafting runtime unchanged. Called only on the parallel_drafting
branch of setup.
Forward+backward one batch via P-EAGLE sequence partitioning.
Builds the segment plan, then runs one DDP.forward per segment and
back-propagates each here (the recipe owns backward() so DDP’s
gradient all-reduce fires). no_sync defers the all-reduce on every
segment except the last, so there is exactly one all-reduce per
micro-batch — matching the single-pass path and keeping the per-rank
collective count aligned. Gradients are divided by
grad_accumulation_steps exactly like the single-pass backward; the
returned (detached) metrics aggregate the whole batch for logging.
Move a batch to device and run the live target to its draft supervision.
Used by the partitioned step; P-EAGLE requires the live target (the offline cache is EAGLE-3 TTT-only).
Construct the P-EAGLE trainer and record the sequence-partitioning flag.
sequence_partitions (S) is a training-only memory knob: when > 1 the
P-EAGLE trainer splits each sequence into S segments and runs a separate
forward+backward per segment, accumulating gradients, so only one
segment’s activations are resident at once (P-EAGLE Algorithm 1,
arXiv:2602.01469). It does not alter the loss or the saved checkpoint, so
it is NOT serialized into the draft config. Default 1 == single flat
forward. num_depths is P-EAGLE’s K (number of parallel COD depths),
default 8 to match speculators. Called only on the parallel_drafting
branch of setup.