nemo_automodel.recipes.llm.train_dflash
nemo_automodel.recipes.llm.train_dflash
DFlash draft-model training recipe (Qwen3-style and Kimi K3 targets).
DFlash drafts a whole block of tokens in parallel via MASK-token denoising
conditioned on the frozen target’s hidden states (see
nemo_automodel.components.speculative.dflash). This recipe mirrors the EAGLE
recipes’ scaffolding — online target hidden-state capture, gradient
accumulation with a trailing-window flush, and the same checkpointer plumbing —
but trains the DFlash draft with its block-wise cross-entropy objective.
Module Contents
Classes
Functions
Data
API
Bases: BaseRecipe
Recipe for DFlash draft-model training on Qwen3-style dense / MoE and Kimi K3 targets.
Build the checkpointer using the same plumbing as the EAGLE recipes.
Build the draft dflash_config block. Subclasses extend it (e.g. Domino).
Derive the draft config for a Qwen3-shaped target.
A small non-causal Qwen3 stack that reuses the target’s architecture defaults (head_dim, rope_theta, rms_norm_eps, …). Targets whose draft is not Qwen3-shaped register their own builder on the spec instead and never reach this.
Parameters:
The recipe’s recipe_args mapping.
The target’s decoder config.
The draft class being built, stamped into architectures.
Depth of the draft stack.
Depth of the target, for the fc input width.
Target layers captured as draft context.
The draft’s attention implementation.
Returns: Qwen3Config
The draft Qwen3Config.
Load the frozen (optionally tensor-parallel) target model.
draft_spec.build_target_kwargs supplies any architecture-specific
from_pretrained arguments (Kimi K3 pins the text-only architecture and
an expert-parallel backend); it is empty for a Qwen3-shaped target.
With a distributed: section and tp_size>1 the target is sharded
in place by from_pretrained (its FSDP2 parallelize plan); the small
draft stays replicated and runs DDP over the “dp” axis (which excludes
“tp”), and the trainer module gathers the target’s vocab-sharded lm_head
/ embed_tokens outputs. Absent, the original single-GPU-per-rank DP path
is used. Sets self.dist_setup / self.device_mesh / self.dp_mesh
as a side effect and returns the (grad-disabled) target.
Build the frozen-target hidden-state capture wrapper.
Subclasses override to capture extra teacher signals (e.g. JetSpec also captures the target logits for its forward-KL distillation).
Build the trainer wrapper. Subclasses override to swap the wrapper (e.g. Domino).
Pick the draft class from the resolved spec.
Subclasses override to select a different draft of the same family; the
DFlash 2 recipe returns draft_spec.draft2_cls. The returned class name
is also what lands in the saved config’s architectures, which is how a
serving engine tells the two drafts apart.
Process group for the draft’s gradient all-reduce.
With tensor parallelism the draft is replicated across tp ranks, so a
full-world all-reduce would average duplicate gradients; restrict it to
the “dp” sub-axis (which excludes tp) so it reduces only across real data
replicas. Without a mesh (tp_size=1) dp_mesh is None -> return None ->
the default full-world group, unchanged.
Create zeroed subclass validation accumulators on the trainer device.
Return additional validation numerator and denominator pairs.
The base DFlash metrics are accumulated directly by _run_eval.
Subclasses use this hook for extra scalar statistics, with both tensors
on the same device as metrics.loss so they can participate in the
same ordered distributed SUM reductions.
Return algorithm-specific training numerator and denominator pairs.
These are accumulated over the micro-batches between two log points and
divided at the log point, the same way train/loss and
train/accuracy are, so every curve on the dashboard covers the same
window. Returning the per-micro-batch mean instead would report a single
micro-batch out of log_every_steps * grad_accumulation_steps.
Restore DFlash meta: global_step and epoch, and validate mask_token_id.
Hook for subclasses to log extra per-step metrics at a log point (no-op here).
Log a saved checkpoint on rank 0 when checkpointing is enabled.
Always save the fully-trained model at the end, unless a cadence already saved the final step.
Save a checkpoint mid-epoch when ckpt_every_steps is configured.
Resolve and validate the MASK token id that fills non-anchor block positions.
DFlash fills every non-anchor slot of a [anchor, MASK, MASK, ...] block
with this id, and the draft’s embed_tokens row at that id becomes the
learned “predict here” signal. It must be chosen deliberately (a reserved /
unused token), exactly like P-EAGLE’s mask_token_id: the previous silent
fallback to tokenizer.pad_token_id was unsafe because pad is commonly
aliased to eos (or another meaningful token), which conflates the mask
signal with real content and quietly degrades acceptance without erroring.
Require it explicitly and range-check it; the inference runtime must fill the
block slots with the same id.
Run one trainer-module forward. Subclasses override to inject extra inputs (e.g. lambda_base).
Persist DFlash meta: global_step, epoch, block_size, and target layers.
Log scalar metrics to the rank-zero W&B run when configured.
Restore the DFlash draft model, optimizer, scheduler, RNG, and global_step.
Run the DFlash training loop.
Persist the DFlash draft model, optimizer, scheduler, RNG, and meta.
Build the target model, DFlash draft, data, optimizer, and trainer module.
Min-reduce a per-rank “this micro-batch has valid anchors” flag.
Under DDP a data-dependent NoValidAnchorsError skip is per-rank: if one
rank skips its backward (and its gradient all-reduce) while another runs its,
the collective mismatches (hang) and the accumulation windows desync. Taking
the MIN across ranks makes the skip decision unanimous — every rank skips the
micro-batch unless all of them have something to learn from it. The reduce is
a tiny independent collective, safe inside no_sync (which only gates the
DDP backward all-reduce). Single-process runs return the local flag unchanged.
Sum a scalar metric tensor across all distributed ranks in place.
Sequence-packing metadata from a dataloader batch (empty dict when unpacked).
Project the target’s decoder config onto the keys a plain Qwen3 config declares.
The draft is always a Qwen3-shaped stack, but its config starts from the
target’s, and a Qwen3.5 text config carries fields the draft has no use for:
linear-attention shapes, MTP, output gating, and partial_rotary_factor
(top-level and inside rope_parameters, on its own key or as mRoPE
sections). None of them may reach the saved draft config — the published
drafters ship without them, and they are not inert there: the HF Qwen3 stack
the draft trains with applies full rotary regardless, so a serving runtime
that honours a leaked partial_rotary_factor: 0.25 would rebuild the
rotary table at a quarter width and silently mismatch the trained weights.
Parameters:
to_dict() of the target’s decoder config.
Returns: dict
The subset of target_text_config a Qwen3Config declares, with
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,
the dataloader sampler, and the checkpointer dp_rank 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 DFlash path cannot honor (fail fast at setup).
Context parallelism shards the sequence and strips the block-causal mask packing
relies on, and a FlashAttention target packs documents from per-document
position_ids only at batch size 1.
Entrypoint for TrainDFlashRecipe.