nemo_automodel.recipes.llm.train_dflash

View as Markdown

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

NameDescription
TrainDFlashRecipeRecipe for DFlash draft-model training on Qwen3-style dense / MoE and Kimi K3 targets.

Functions

NameDescription
_all_ranks_have_validMin-reduce a per-rank “this micro-batch has valid anchors” flag.
_all_reduce_sumSum a scalar metric tensor across all distributed ranks in place.
_packing_kwargsSequence-packing metadata from a dataloader batch (empty dict when unpacked).
_project_onto_qwen3_config_keysProject the target’s decoder config onto the keys a plain Qwen3 config declares.
_submesh_or_noneReturn the named (flattened) submesh, or None if absent / no mesh.
_validate_packing_gatesReject sequence-packing configs the DFlash path cannot honor (fail fast at setup).
mainEntrypoint for TrainDFlashRecipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe(
cfg
)

Bases: BaseRecipe

Recipe for DFlash draft-model training on Qwen3-style dense / MoE and Kimi K3 targets.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_checkpointer(
target_path: str
) -> None

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

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_dflash_config(
recipe_cfg,
target_layer_ids: list[int]
) -> dict

Build the draft dflash_config block. Subclasses extend it (e.g. Domino).

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_qwen3_draft_config(
recipe_cfg,
target_text_config,
draft_cls,
draft_num_hidden_layers: int,
num_target_layers: int,
target_layer_ids: list[int],
attention_backend: str
) -> transformers.models.qwen3.configuration_qwen3.Qwen3Config

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:

recipe_cfg

The recipe’s recipe_args mapping.

target_text_config

The target’s decoder config.

draft_cls

The draft class being built, stamped into architectures.

draft_num_hidden_layers
int

Depth of the draft stack.

num_target_layers
int

Depth of the target, for the fc input width.

target_layer_ids
list[int]

Target layers captured as draft context.

attention_backend
str

The draft’s attention implementation.

Returns: Qwen3Config

The draft Qwen3Config.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_target_model(
recipe_cfg,
target_path: str,
) -> torch.nn.Module

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.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_target_wrapper(
target_layer_ids: list[int]

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).

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_trainer_module(
attention_backend: str,
recipe_cfg
)

Build the trainer wrapper. Subclasses override to swap the wrapper (e.g. Domino).

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._draft_cls(
) -> type[torch.nn.Module]

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.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._draft_ddp_process_group()

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.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._empty_extra_eval_metric_sums() -> dict[str, list[torch.Tensor]]

Create zeroed subclass validation accumulators on the trainer device.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._extra_eval_metric_sums(
metrics
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]

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.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._extra_train_metric_sums(
metrics
) -> dict[str, tuple[float, float]]

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.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._load_extra_state(
ckpt_dir: str
) -> None

Restore DFlash meta: global_step and epoch, and validate mask_token_id.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._log_extra_train_metrics(
epoch_idx: int
) -> None

Hook for subclasses to log extra per-step metrics at a log point (no-op here).

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._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_dflash.TrainDFlashRecipe._maybe_save_final_checkpoint(
completed_epochs: int
) -> bool

Always save the fully-trained model at the end, unless a cadence already saved the final step.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._maybe_save_step_checkpoint(
epoch: int
) -> bool

Save a checkpoint mid-epoch when ckpt_every_steps is configured.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._module()
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._resolve_mask_token_id(
recipe_cfg,
vocab_size: int
) -> int
staticmethod

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.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._run_eval()
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._run_trainer_step(
target_batch
)

Run one trainer-module forward. Subclasses override to inject extra inputs (e.g. lambda_base).

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._save_extra_state(
path: str,
epoch: int
) -> None

Persist DFlash meta: global_step, epoch, block_size, and target layers.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._wandb_log(
data: dict[str, float],
step: int
) -> None

Log scalar metrics to the rank-zero W&B run when configured.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe.load_checkpoint(
restore_from: str | None = None
) -> None

Restore the DFlash draft model, optimizer, scheduler, RNG, and global_step.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe.run_train_validation_loop()

Run the DFlash training loop.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe.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 the DFlash draft model, optimizer, scheduler, RNG, and meta.

nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe.setup()

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

nemo_automodel.recipes.llm.train_dflash._all_ranks_have_valid(
local_has_valid: int,
is_ddp: bool,
device
) -> bool

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.

nemo_automodel.recipes.llm.train_dflash._all_reduce_sum(
value: torch.Tensor
) -> torch.Tensor

Sum a scalar metric tensor across all distributed ranks in place.

nemo_automodel.recipes.llm.train_dflash._packing_kwargs(
batch: dict[str, torch.Tensor]
) -> dict[str, torch.Tensor]

Sequence-packing metadata from a dataloader batch (empty dict when unpacked).

nemo_automodel.recipes.llm.train_dflash._project_onto_qwen3_config_keys(
target_text_config: dict
) -> dict

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:

target_text_config
dict

to_dict() of the target’s decoder config.

Returns: dict

The subset of target_text_config a Qwen3Config declares, with

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

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).

nemo_automodel.recipes.llm.train_dflash._validate_packing_gates(
cp_size: int,
target_attn_impl: str,
micro_batch_size: int
) -> None

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.

nemo_automodel.recipes.llm.train_dflash.main(
config_path: str | None = None
)

Entrypoint for TrainDFlashRecipe.

nemo_automodel.recipes.llm.train_dflash.logger = logging.getLogger(__name__)