nemo_automodel.recipes.llm.train_dflash2

View as Markdown

DFlash 2 draft-model training recipe (Qwen3-style targets).

DFlash 2 (https://inco.ai/blog/dflash2/) keeps DFlash’s one-pass parallel block draft and adds two cheap modules: a two-tap dynamic convolution around every draft sublayer, which carries the short-range within-block work and removes most of DFlash’s suffix decay, and a pairwise path selector that walks one coherent path through each position’s top-k candidates instead of keeping every top-1 pick independently. See nemo_automodel.components.speculative.dflash.dflash2_core.

This recipe reuses every piece of the DFlash recipe — online target hidden-state capture, anchor sampling, the block attention mask, gradient accumulation, and checkpointing — and only swaps in the DFlash 2 draft class, stamps the extra dflash_config fields a serving engine needs to rebuild it, and swaps in the DFlash 2 trainer wrapper.

IMPORTANT: as with DFlash, regenerate the training responses with the target model first — training teacher-forces ground-truth tokens while inference is autoregressive, and the distribution mismatch hurts acceptance length otherwise.

Module Contents

Classes

NameDescription
TrainDFlash2RecipeRecipe for DFlash 2 draft-model training: in-block convolutions + path selector.

Functions

NameDescription
mainEntrypoint for TrainDFlash2Recipe.

Data

logger

API

class nemo_automodel.recipes.llm.train_dflash2.TrainDFlash2Recipe()

Bases: TrainDFlashRecipe

Recipe for DFlash 2 draft-model training: in-block convolutions + path selector.

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

Extend the DFlash draft config with the convolution and selector shapes.

nemo_automodel.recipes.llm.train_dflash2.TrainDFlash2Recipe._build_trainer_module(
attention_backend: str,
recipe_cfg
)

Build the DFlash 2 trainer wrapper (block CE + candidate-selection CE).

nemo_automodel.recipes.llm.train_dflash2.TrainDFlash2Recipe._draft_cls(
) -> type[torch.nn.Module]

Build the DFlash 2 draft instead of the plain DFlash one.

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

Create rank-symmetric DFlash 2 validation accumulators.

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

Return additive DFlash 2 validation statistics.

Every returned value is a pair of scalar tensors on the trainer device. The shared validation loop SUM-reduces each pair before division.

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

Return the selector and backbone diagnostics as window sums.

The two loss terms are per-micro-batch means, so their denominator is the micro-batch count; the accuracy-like curves are token- or block-weighted, matching how train/loss and train/accuracy are averaged.

nemo_automodel.recipes.llm.train_dflash2.TrainDFlash2Recipe._log_extra_train_metrics(
epoch_idx: int
) -> None

Log the DFlash 2 diagnostics for the most recent step (rank-0 local).

nemo_automodel.recipes.llm.train_dflash2.TrainDFlash2Recipe._run_trainer_step(
target_batch
)

Forward through the DFlash 2 wrapper and cache the step’s diagnostics.

nemo_automodel.recipes.llm.train_dflash2.TrainDFlash2Recipe.setup()

Build everything via the DFlash recipe, then reset the per-step metric cache.

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

Entrypoint for TrainDFlash2Recipe.

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