nemo_automodel.recipes.dllm.strategy
nemo_automodel.recipes.dllm.strategy
Model-specific strategies for diffusion LLM (dLLM) training.
Each strategy encapsulates the variation points that differ across dLLM model families:
- Loss function creation — which loss module to use.
- Pre-step processing — corruption (MDLM) or target-model forwards (DFlash).
- Forward-backward — the per-microbatch forward + loss + backward.
- Normalization mode — loss denominator: supervised tokens or noise tokens.
- Extra setup — loading auxiliary models (e.g. frozen target for DFlash).
To add a new dLLM variant, implement a DLLMStrategy subclass and
register it in DLLM_STRATEGIES. No changes to the recipe are required.
Module Contents
Classes
Functions
Data
API
Bases: DLLMStrategy
Strategy for diffusion_gemma block-diffusion SFT (single-turn v1).
- Loss:
BlockDiffusionCrossEntropyLoss(flat CE, no1/p, no AR). - Corruption:
corrupt_uniform_random— per-blockt~U(eps,1), supervised positions replaced with uniform random vocab tokens (no[MASK]). Requiresvocab_size(seecreate_loss_fn). - Normalization:
"supervised"(denominator = all supervised canvas-token count, matching Google’s all-canvas loss support — NOT corrupted-only). - Batch: the encoder sees the clean full sequence (prompt + response);
the decoder canvas is the noised response region only, sliced and the
block-causal mask built by
DiffusionGemmaSFTRecipe(the recipe owns the response-window construction because it also needs the sliced loss tensors).prepare_batchonly assigns the encoder input and the full noised sequence;split_prompt_responsegives the per-example prompt boundary the recipe slices on.
v1 is single-turn only: multi-turn ChatDataset loss_mask is
0..0 1..1 0..0 1..1 (not a contiguous suffix), so the prompt|response
split is ill-defined. Interleaved multi-turn masking is deferred.
Single-turn boundary: prefix length(s) and a per-position response mask.
Returns (prefix_lengths, response_mask) where prefix_lengths[b] is
the index of the first supervised (loss_mask == 1) position in row
b (the start of the response) and response_mask is the boolean
mask of response positions (position >= prefix_length). Used by the
_forward_backward_step override to build canvas_ids and the
block-causal mask’s prefix_lengths.
Single-turn assumption: loss_mask is a single contiguous suffix.
Bases: DLLMStrategy
Strategy for DFlash dual-model draft training.
DFlash training differs from MDLM in three ways:
- A frozen causal target LM provides hidden-state context.
- One clean anchor token starts each block; the rest are mask-filled.
- Loss is decay-weighted by position within the block (Eq. 4).
All DFlash-specific logic lives here so DiffusionLMSFTRecipe
requires no subclassing for DFlash.
YAML configuration (under the dflash: key):
target_model_id(required) — frozen causal LM hub ID.target_torch_dtype(default"bfloat16") — target dtype string.block_size(default 0) — draft block size; 0 reads from draft config.loss_decay_gamma(default 0.0) — γ for Eq. 4; 0 uses paper defaults.num_blocks_per_sample(default 1) — N anchor blocks per sequence per step, enabling the multi-block sparse-attention pass from §4.2. Paper default is 512 (Appendix A.1); requiresattention_backend=flex_attention.attention_backend(default"sdpa") —"sdpa"materialises a dense[B, 1, N·bs, S+N·bs]mask (OOMs at high N);"flex_attention"uses a sparseBlockMaskand matches the paper’s setup.overlap_anchors(defaultTrue) — whenTrue, anchors are sampled independently (paper behaviour); whenFalse, anchors are forced non-overlapping (stars-and-bars, caps atseq_len // block_size).
Sample num_blocks anchors per sample and gather block tensors.
Each sequence in the batch independently draws N = num_blocks anchor
positions from its own [1, valid_len_b - block_size] range (paper
§4.2 “randomly sample anchor tokens”). Per-sample sampling gives more
position diversity per step than sharing one anchor set across the batch.
Samples with fewer than 2 * block_size supervised tokens are dropped
via block_keep_mask so degenerate short sequences contribute no loss.
Returns: torch.Tensor
[B, N] long — per-sample anchor positions.
DFlash microbatch: draft forward + decay loss + (optional) backward.
Sample anchor blocks and run frozen target forwards for all microbatches.
Load and freeze the target LM; resolve block_size, layer_ids, decay loss.
Abstract base for dLLM model strategies.
Metric key used for dLLM loss in MetricsSample and console log lines.
Token count used as the loss denominator: "supervised" or "noise".
"supervised"— totalloss_mask == 1positions (default)."noise"— actually-corrupted positions (noise_mask == True).
Return (noisy_input_ids, noise_mask, p_mask).
generator (optional): a step-seeded torch.Generator for the
corruption draws so noise reproduces on resume. Strategies that draw from
the global RNG may ignore it; block_diffusion threads it through.
Return the loss module for this model type.
Run one microbatch forward + loss + (optionally) backward.
Default implementation delegates to the recipe’s existing MDLM
_forward_backward_step so that the MDLM code path is unchanged.
Pre-process all microbatches before the forward-backward loop.
Called once per training step (and once per val batch) with the full
list of microbatch dicts. May mutate batch dicts in-place to stash
pre-computed tensors for forward_backward.
Returns: int
(num_noise_tokens, num_supervised_tokens) — raw (un-allreduced)
Mutate batch in-place for the model’s forward pass and return it.
Hook called at the end of DiffusionLMSFTRecipe.setup.
Strategies that need auxiliary models (e.g. a frozen target LM) or
that resolve recipe.mask_token_id should do so here.
Bases: DLLMStrategy
Strategy for hybrid diffusion + AR models (e.g., Nemotron-Labs-Diffusion).
- Loss:
HybridDiffusionLLMLosswith configurablear_loss_alpha. - Corruption: uniform when
block_sizeisNone, blockwise otherwise. - Batch: model receives clean tokens +
masked_indicessidecar; the model applies masking internally during its forward pass. - Normalization: hybrid models normalize diffusion loss by the corrupted (noise) token count, not the full supervised count.
Bases: DLLMStrategy
Strategy for Introspective Diffusion LM (I-DLM) all-masked finetuning.
Converts an AR causal LM into a diffusion LM (Yu et al., 2026):
- Corruption: deterministic all-masked over the supervised region
(
corrupt_all_masked). - Forward: the noisy and clean copies are concatenated into a length-
2L[x_t | x_0]sequence and run under the block-diffusion attention mask (create_idlm_sdpa_mask/create_idlm_block_mask), so decode tokens attend the clean ground-truth prefix and the clean copy stays strict-causal. - Loss:
IDLMLoss(Dream-shiftedCE_noisy + alpha*CE_clean, both supervised on the response).
The mask is built for sdpa/eager (dense additive) or
flex_attention (sparse BlockMask, preferred at scale); FlashAttention-2
is unsupported (it ignores arbitrary masks), as is context parallelism.
I-DLM microbatch: single [x_t | x_0] forward + two-CE loss.
Required by the abstract interface but unused on the I-DLM path.
forward_backward is overridden and builds the [x_t | x_0]
concat itself from the pre_step sidecars, so the recipe never routes
an I-DLM batch through here. Kept (and kept correct) only to satisfy
DLLMStrategy; assigning the noisy ids matches what the base
MDLM path would do if a future edit re-enabled that route.
Bases: DLLMStrategy
Strategy for MDLM / LLaDA-style models.
- Loss:
MDLMCrossEntropyLoss - Corruption: uniform masking (
corrupt_uniform) - Batch: model receives noisy (corrupted) tokens as
input_ids
Bases: DLLMStrategy
Strategy for SCDD (Self-Correcting Discrete Diffusion).
Paper: https://openreview.net/forum?id=zQKlzKB6I9
SCDD generalises MDLM by adding a uniform-transition channel to the absorbing forward process, so the denoiser is trained on contexts that contain wrong-but-plausible tokens and learns to overwrite them. That self-correction is what lets it decode many tokens per step without the quality collapse a pure absorbing model shows under parallel decoding.
- Loss:
SCDDLoss— the discrete-time NELBO with a denoising term at[MASK]positions and a correction term everywhere else. - Corruption:
corrupt_mixdriven byscdd_scheduleat a diffusion time drawn on the1/Tgrid. - Normalization:
"supervised"— the ELBO is supported on every supervised position, not only the corrupted ones. - Batch: like MDLM, the model receives the corrupted tokens as
input_idsand attends bidirectionally.
Time conditioning: the SCDD reference backbone takes the noise level as an
input. Pretrained masked-dLLM checkpoints in Automodel (LLaDA and friends)
are time-free — they read the corruption level off the number of visible
[MASK] tokens — so no time embedding is threaded into the forward pass
here, matching MDLMStrategy. The schedule still enters the
objective through the ELBO weights.
Requires dllm.vocab_size and dllm.mask_token_id: the uniform channel
samples replacements over the vocabulary minus [MASK], and the loss
re-parameterises the model output over that same domain. Context parallelism
is unsupported — the ELBO scores the corrupted tokens against the clean
targets, which the recipe keeps unsharded.
Evenly-spaced target hidden-layer indices for DFlash feature extraction.
Look up and instantiate a dLLM strategy by mode name.
Raises:
ValueError: If mode is not registered inDLLM_STRATEGIES.