nemo_automodel.recipes.llm.train_dflash2
nemo_automodel.recipes.llm.train_dflash2
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
Functions
Data
API
Bases: TrainDFlashRecipe
Recipe for DFlash 2 draft-model training: in-block convolutions + path selector.
Extend the DFlash draft config with the convolution and selector shapes.
Build the DFlash 2 trainer wrapper (block CE + candidate-selection CE).
Build the DFlash 2 draft instead of the plain DFlash one.
Create rank-symmetric DFlash 2 validation accumulators.
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.
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.
Log the DFlash 2 diagnostics for the most recent step (rank-0 local).
Forward through the DFlash 2 wrapper and cache the step’s diagnostics.
Build everything via the DFlash recipe, then reset the per-step metric cache.
Entrypoint for TrainDFlash2Recipe.