> This page is for version Nightly (default).
> For other versions, use one of these documentation indexes:
> - Nightly (default): https://docs.nvidia.com/nemo/automodel/nightly/llms.txt
> - Latest: https://docs.nvidia.com/nemo/automodel/latest/llms.txt
> - 0.5.0 · 26.06: https://docs.nvidia.com/nemo/automodel/v0.5/llms.txt
> - 0.4.0 · 26.04: https://docs.nvidia.com/nemo/automodel/v0.4/llms.txt

> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# 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

| Name                                                                              | Description                                                                            |
| --------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------- |
| [`TrainDFlashRecipe`](#nemo_automodel-recipes-llm-train_dflash-TrainDFlashRecipe) | Recipe for DFlash draft-model training on Qwen3-style dense / MoE and Kimi K3 targets. |

### Functions

| Name                                                                                                          | Description                                                                        |
| ------------------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------- |
| [`_all_ranks_have_valid`](#nemo_automodel-recipes-llm-train_dflash-_all_ranks_have_valid)                     | Min-reduce a per-rank "this micro-batch has valid anchors" flag.                   |
| [`_all_reduce_sum`](#nemo_automodel-recipes-llm-train_dflash-_all_reduce_sum)                                 | Sum a scalar metric tensor across all distributed ranks in place.                  |
| [`_packing_kwargs`](#nemo_automodel-recipes-llm-train_dflash-_packing_kwargs)                                 | Sequence-packing metadata from a dataloader batch (empty dict when unpacked).      |
| [`_project_onto_qwen3_config_keys`](#nemo_automodel-recipes-llm-train_dflash-_project_onto_qwen3_config_keys) | Project the target's decoder config onto the keys a plain Qwen3 config declares.   |
| [`_submesh_or_none`](#nemo_automodel-recipes-llm-train_dflash-_submesh_or_none)                               | Return the named (flattened) submesh, or None if absent / no mesh.                 |
| [`_validate_packing_gates`](#nemo_automodel-recipes-llm-train_dflash-_validate_packing_gates)                 | Reject sequence-packing configs the DFlash path cannot honor (fail fast at setup). |
| [`main`](#nemo_automodel-recipes-llm-train_dflash-main)                                                       | Entrypoint for `TrainDFlashRecipe`.                                                |

### Data

[`logger`](#nemo_automodel-recipes-llm-train_dflash-logger)

### API

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

**Bases:** [BaseRecipe](/nemo-automodel/nemo_automodel/recipes/base_recipe#nemo_automodel-recipes-base_recipe-BaseRecipe)

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

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

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

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

```python
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`.

```python
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_target_model(
    recipe_cfg,
    target_path: str,
    draft_spec: nemo_automodel.components.speculative.dflash.registry.DFlashDraftSpec
) -> 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&gt;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.

```python
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._build_target_wrapper(
    target_layer_ids: list[int]
) -> nemo_automodel.components.speculative.dflash.target.HFDFlashTargetModel
```

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

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

```python
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._draft_cls(
    draft_spec: nemo_automodel.components.speculative.dflash.registry.DFlashDraftSpec
) -> 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.

```python
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.

```python
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.

```python
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.

```python
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`.

```python
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.

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

```python
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.

```python
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.

```python
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.

```python
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._module()
```

```python
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.

```python
nemo_automodel.recipes.llm.train_dflash.TrainDFlashRecipe._run_eval()
```

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

```python
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.

```python
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.

```python
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.

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

Run the DFlash training loop.

```python
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.

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

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

```python
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.

```python
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.

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

```python
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

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

```python
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.

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

Entrypoint for `TrainDFlashRecipe`.

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