nemo_automodel.recipes.dllm.strategy

View as Markdown

Model-specific strategies for diffusion LLM (dLLM) training.

Each strategy encapsulates the variation points that differ across dLLM model families:

  1. Loss function creation — which loss module to use.
  2. Pre-step processing — corruption (MDLM) or target-model forwards (DFlash).
  3. Forward-backward — the per-microbatch forward + loss + backward.
  4. Normalization mode — loss denominator: supervised tokens or noise tokens.
  5. 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

NameDescription
BlockDiffusionStrategyStrategy for diffusion_gemma block-diffusion SFT (single-turn v1).
DFlashStrategyStrategy for DFlash dual-model draft training.
DLLMStrategyAbstract base for dLLM model strategies.
HybridStrategyStrategy for hybrid diffusion + AR models (e.g., Nemotron-Labs-Diffusion).
IDLMStrategyStrategy for Introspective Diffusion LM (I-DLM) all-masked finetuning.
MDLMStrategyStrategy for MDLM / LLaDA-style models.
SCDDStrategyStrategy for SCDD (Self-Correcting Discrete Diffusion).

Functions

NameDescription
_build_target_layer_idsEvenly-spaced target hidden-layer indices for DFlash feature extraction.
get_dllm_strategyLook up and instantiate a dLLM strategy by mode name.

Data

DLLM_STRATEGIES

logger

API

class nemo_automodel.recipes.dllm.strategy.BlockDiffusionStrategy()

Bases: DLLMStrategy

Strategy for diffusion_gemma block-diffusion SFT (single-turn v1).

  • Loss: BlockDiffusionCrossEntropyLoss (flat CE, no 1/p, no AR).
  • Corruption: corrupt_uniform_random — per-block t~U(eps,1), supervised positions replaced with uniform random vocab tokens (no [MASK]). Requires vocab_size (see create_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_batch only assigns the encoder input and the full noised sequence; split_prompt_response gives 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.

_vocab_size
int | None = None
normalization_mode
str
nemo_automodel.recipes.dllm.strategy.BlockDiffusionStrategy.apply_corruption(
input_ids,
loss_mask,
mask_token_id,
eps,
block_size,
half_life_ratio,
generator = None
)
nemo_automodel.recipes.dllm.strategy.BlockDiffusionStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
nemo_automodel.recipes.dllm.strategy.BlockDiffusionStrategy.prepare_batch(
batch,
noisy_input_ids,
noise_mask,
clean_input_ids
)
nemo_automodel.recipes.dllm.strategy.BlockDiffusionStrategy.split_prompt_response(
input_ids: torch.Tensor,
loss_mask: torch.Tensor
) -> typing.Tuple[torch.Tensor, torch.Tensor]
staticmethod

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.

class nemo_automodel.recipes.dllm.strategy.DFlashStrategy()

Bases: DLLMStrategy

Strategy for DFlash dual-model draft training.

DFlash training differs from MDLM in three ways:

  1. A frozen causal target LM provides hidden-state context.
  2. One clean anchor token starts each block; the rest are mask-filled.
  3. 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); requires attention_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 sparse BlockMask and matches the paper’s setup.
  • overlap_anchors (default True) — when True, anchors are sampled independently (paper behaviour); when False, anchors are forced non-overlapping (stars-and-bars, caps at seq_len // block_size).
attention_backend
str = 'sdpa'
block_size
int = 0
fixed_ctx_len
int = 0
layer_ids
list = []
loss_log_key
str
num_blocks_per_sample
int = 1
overlap_anchors
bool = True
use_fused_linear_ce
bool = True
nemo_automodel.recipes.dllm.strategy.DFlashStrategy._run_target_forward(
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
start: int
) -> torch.Tensor
nemo_automodel.recipes.dllm.strategy.DFlashStrategy._sample_anchor_block(
recipe,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
loss_mask: torch.Tensor | None = None
) -> tuple[int, torch.Tensor, torch.Tensor, torch.Tensor]
nemo_automodel.recipes.dllm.strategy.DFlashStrategy._sample_anchor_blocks(
recipe,
input_ids: torch.Tensor,
attn: torch.Tensor,
num_blocks: int,
loss_mask: torch.Tensor | None = None
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]

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.

nemo_automodel.recipes.dllm.strategy.DFlashStrategy.apply_corruption(
input_ids,
loss_mask,
mask_token_id,
eps,
block_size,
half_life_ratio,
generator = None
)
nemo_automodel.recipes.dllm.strategy.DFlashStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
nemo_automodel.recipes.dllm.strategy.DFlashStrategy.forward_backward(
recipe,
idx: int,
batch: dict,
loss_buffer: list,
num_diffusion_tokens: int,
num_ar_tokens: int | None = None,
num_batches: int,
is_train: bool = True
) -> None

DFlash microbatch: draft forward + decay loss + (optional) backward.

nemo_automodel.recipes.dllm.strategy.DFlashStrategy.pre_step(
recipe,
batches
) -> tuple[int, int]

Sample anchor blocks and run frozen target forwards for all microbatches.

nemo_automodel.recipes.dllm.strategy.DFlashStrategy.prepare_batch(
batch,
noisy_input_ids,
noise_mask,
clean_input_ids
)
nemo_automodel.recipes.dllm.strategy.DFlashStrategy.setup_extra(
recipe
) -> None

Load and freeze the target LM; resolve block_size, layer_ids, decay loss.

class nemo_automodel.recipes.dllm.strategy.DLLMStrategy()
Abstract

Abstract base for dLLM model strategies.

loss_log_key
str

Metric key used for dLLM loss in MetricsSample and console log lines.

normalization_mode
str

Token count used as the loss denominator: "supervised" or "noise".

  • "supervised" — total loss_mask == 1 positions (default).
  • "noise" — actually-corrupted positions (noise_mask == True).
nemo_automodel.recipes.dllm.strategy.DLLMStrategy.apply_corruption(
input_ids: torch.Tensor,
loss_mask: torch.Tensor,
mask_token_id: int,
eps: float,
block_size: int | None,
half_life_ratio: float | None,
generator: torch.Generator | None = None
) -> typing.Tuple[torch.Tensor, torch.Tensor, torch.Tensor]
abstract

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.

nemo_automodel.recipes.dllm.strategy.DLLMStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
abstract

Return the loss module for this model type.

nemo_automodel.recipes.dllm.strategy.DLLMStrategy.forward_backward(
recipe,
idx: int,
batch: dict,
loss_buffer: list,
num_diffusion_tokens: int,
num_ar_tokens: int | None = None,
num_batches: int,
is_train: bool = True
) -> None

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.

nemo_automodel.recipes.dllm.strategy.DLLMStrategy.pre_step(
recipe,
batches
) -> tuple[int, int]

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)

nemo_automodel.recipes.dllm.strategy.DLLMStrategy.prepare_batch(
batch: typing.Dict[str, torch.Tensor],
noisy_input_ids: torch.Tensor,
noise_mask: torch.Tensor,
clean_input_ids: torch.Tensor
) -> typing.Dict[str, torch.Tensor]
abstract

Mutate batch in-place for the model’s forward pass and return it.

nemo_automodel.recipes.dllm.strategy.DLLMStrategy.setup_extra(
recipe
) -> None

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.

class nemo_automodel.recipes.dllm.strategy.HybridStrategy()

Bases: DLLMStrategy

Strategy for hybrid diffusion + AR models (e.g., Nemotron-Labs-Diffusion).

  • Loss: HybridDiffusionLLMLoss with configurable ar_loss_alpha.
  • Corruption: uniform when block_size is None, blockwise otherwise.
  • Batch: model receives clean tokens + masked_indices sidecar; 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.
normalization_mode
str
nemo_automodel.recipes.dllm.strategy.HybridStrategy.apply_corruption(
input_ids,
loss_mask,
mask_token_id,
eps,
block_size,
half_life_ratio,
generator = None
)
nemo_automodel.recipes.dllm.strategy.HybridStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
nemo_automodel.recipes.dllm.strategy.HybridStrategy.prepare_batch(
batch,
noisy_input_ids,
noise_mask,
clean_input_ids
)
class nemo_automodel.recipes.dllm.strategy.IDLMStrategy()

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-shifted CE_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.

block_size
= 1
nemo_automodel.recipes.dllm.strategy.IDLMStrategy.apply_corruption(
input_ids,
loss_mask,
mask_token_id,
eps,
block_size,
half_life_ratio,
generator = None
)
nemo_automodel.recipes.dllm.strategy.IDLMStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
nemo_automodel.recipes.dllm.strategy.IDLMStrategy.forward_backward(
recipe,
idx: int,
batch: dict,
loss_buffer: list,
num_diffusion_tokens: int,
num_ar_tokens: int | None = None,
num_batches: int,
is_train: bool = True
) -> None

I-DLM microbatch: single [x_t | x_0] forward + two-CE loss.

nemo_automodel.recipes.dllm.strategy.IDLMStrategy.prepare_batch(
batch,
noisy_input_ids,
noise_mask,
clean_input_ids
)

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.

nemo_automodel.recipes.dllm.strategy.IDLMStrategy.setup_extra(
recipe
) -> None
class nemo_automodel.recipes.dllm.strategy.MDLMStrategy()

Bases: DLLMStrategy

Strategy for MDLM / LLaDA-style models.

  • Loss: MDLMCrossEntropyLoss
  • Corruption: uniform masking (corrupt_uniform)
  • Batch: model receives noisy (corrupted) tokens as input_ids
nemo_automodel.recipes.dllm.strategy.MDLMStrategy.apply_corruption(
input_ids,
loss_mask,
mask_token_id,
eps,
block_size,
half_life_ratio,
generator = None
)
nemo_automodel.recipes.dllm.strategy.MDLMStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
nemo_automodel.recipes.dllm.strategy.MDLMStrategy.prepare_batch(
batch,
noisy_input_ids,
noise_mask,
clean_input_ids
)
class nemo_automodel.recipes.dllm.strategy.SCDDStrategy()

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_mix driven by scdd_schedule at a diffusion time drawn on the 1/T grid.
  • 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_ids and 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.

_gamma_shape
float = 1.0
_max_ratio
float = 0.1
_num_timesteps
int = 1000
_t_peak
float = 0.5
_vocab_size
int | None = None
nemo_automodel.recipes.dllm.strategy.SCDDStrategy.apply_corruption(
input_ids,
loss_mask,
mask_token_id,
eps,
block_size,
half_life_ratio,
generator = None
)
nemo_automodel.recipes.dllm.strategy.SCDDStrategy.create_loss_fn(
dllm_cfg: dict
) -> torch.nn.Module
nemo_automodel.recipes.dllm.strategy.SCDDStrategy.prepare_batch(
batch,
noisy_input_ids,
noise_mask,
clean_input_ids
)
nemo_automodel.recipes.dllm.strategy.SCDDStrategy.setup_extra(
recipe
) -> None
nemo_automodel.recipes.dllm.strategy._build_target_layer_ids(
num_target_layers: int,
num_draft_layers: int
) -> list[int]

Evenly-spaced target hidden-layer indices for DFlash feature extraction.

nemo_automodel.recipes.dllm.strategy.get_dllm_strategy(
mode: str

Look up and instantiate a dLLM strategy by mode name.

Raises:

  • ValueError: If mode is not registered in DLLM_STRATEGIES.
nemo_automodel.recipes.dllm.strategy.DLLM_STRATEGIES: Dict[str, type] = {'mdlm': MDLMStrategy, 'scdd': SCDDStrategy, 'hybrid': HybridStrategy, 'idlm': I...
nemo_automodel.recipes.dllm.strategy.logger = logging.getLogger(__name__)