nemo_automodel.recipes.dllm.train_ft

View as Markdown

Diffusion LLM (dLLM) SFT recipe for Automodel.

Extends TrainFinetuneRecipeForNextTokenPrediction to support diffusion LLM training. Instead of next-token prediction, the model is trained as a denoiser: tokens are randomly corrupted and the model predicts the clean token at each position. Loss is weighted by the inverse corruption probability.

Model-specific behaviour (loss function, corruption strategy, batch preparation) is encapsulated in strategy so that new dLLM variants can be added without modifying this recipe. Current modes:

  • mdlm: Pure masked denoising. Uses MDLMCrossEntropyLoss.

Usage::

python -m torch.distributed.run —nproc-per-node=8
nemo_automodel/recipes/dllm/train_ft.py
-c examples/dllm_sft/llada_sft.yaml

Module Contents

Classes

NameDescription
DiffusionGemmaSFTRecipedLLM SFT recipe for the diffusion_gemma block-diffusion model.
DiffusionLMSFTRecipeRecipe for dLLM (diffusion LLM) supervised fine-tuning.

Functions

NameDescription
corruption_sample_seedTopology-independent per-sample corruption seed.
mainMain entry point for dLLM SFT recipe.

Data

_CORRUPTION_SEED_MULT

_CORRUPTION_STREAM_OFFSET

_SEED_MASK

logger

API

class nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe()

Bases: DiffusionLMSFTRecipe

dLLM SFT recipe for the diffusion_gemma block-diffusion model.

Extends DiffusionLMSFTRecipe (dllm.mode = block_diffusion) with the model-specific response-window canvas wiring (single-turn v1):

  • The encoder sees the clean full sequence (prompt + response). The decoder canvas is the noised response region only: per example the contiguous supervised suffix [prefix_len, S) is left-aligned into a [B, R] canvas (R = S - min(prefix_len); shorter responses are right-padded). Canvas position 0 is the first response token, so the diffusion block boundaries align to the response — block i is bidirectional within itself and attends (offset-block-causal, strict >) the clean prompt + clean earlier response blocks in the encoder KV. This matches the reference inference contract (each block conditioned on the prompt + already-generated blocks).
  • The block-causal masks (full + sliding) are built by build_block_diffusion_training_mask with per-example prefix_lengths = prompt length, response_length = response length, enc_len = full sequence length; decoder_position_ids are the response tokens’ absolute positions so their query RoPE aligns with the encoder key RoPE.
  • Because the model returns canvas-only logits ([B, R, V]), the loss tensors (target_ids / noise_mask / loss_mask / p_mask) are sliced to the same response window in _forward_backward_step so they align with the logits (the inherited path passes full-sequence targets against canvas-only logits → misalignment).
  • The two-pass self-conditioning is orchestrated inside model.forward, so the recipe still calls model(**batch) once.

Single-turn assumption. The response is taken to be the single contiguous supervised suffix (BlockDiffusionStrategy.split_prompt_response). Multi-turn ChatDataset loss_mask (0..0 1..1 0..0 1..1) is not a contiguous suffix, so the prompt|response split is ill-defined; interleaved multi-turn masking is deferred.

corrupt_uniform_random does not use a mask_token_id (random-token corruption); this recipe injects a harmless default so the parent setup, which expects one, does not fail.

_response_window_collation
bool = True
nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe._build_batched_block_mask(
prefix_lengths,
response_lengths,
canvas_len,
enc_len,
block_size,
sliding_window,
device,
dtype
)
staticmethod

Assemble the [B, 1, R, enc_len + R] block-causal mask from per-example builds.

Each example satisfies build_block_diffusion_training_mask’s prefix + response_length == enc_len contract exactly, so the builder is called once per example (batch is small) and the result is placed into a padded additive mask. Pad query rows (j >= response_length) are kept well-defined by leaving their canvas self-diagonal unmasked, so the shared transformer never sees an all--inf softmax row; those rows are unsupervised and discarded from the loss.

nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe._build_response_window(
clean_input_ids,
noisy_input_ids,
noise_mask,
loss_mask,
p_mask,
attention_mask = None,
microbatch_idx: int = 0
)

Slice the batch to the response window and build the block-causal masks.

Returns a dict with the canvas-region tensors (all [B, R]) and the forward inputs (canvas ids, masks, position ids). R is the longest response in the batch; shorter responses are right-padded and the pad positions are dropped from the (sliced) loss_mask / noise_mask so they never contribute to the loss.

nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe._compute_loss_denominators(
batches,
num_noise_tokens,
num_supervised_tokens
)

Count the GLOBAL denominators from the FINAL response-window masks.

The raw pre-mask counts over-count: _build_response_window restricts the diffusion loss to ONE step-selected canvas block (one-canvas scheme) and drops padding, while the AR loss is scored on the padding-masked, next-token-shifted supervised mask. Dividing by num_noise_tokens / num_supervised_tokens would include tokens that cannot contribute, under-scaling diffusion (~by the response-block count) and AR (by padding + the shift). We build each microbatch’s window once here (caching it on the batch so the forward reuses the same step-seeded block selection) and sum the actually-scored tokens.

nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe._decide_self_conditioning(
batch_size: int,
microbatch_idx: int = 0
) -> torch.Tensor

Per-EXAMPLE two-pass self-conditioning coins -> [B] bool tensor.

Google draws the coin PER EXAMPLE (uniform(B) < p) and mixes conditioned / zero-conditioned examples within a batch. Returns a [B] mask so the model gates the self-cond branch per example, while pass-1 (the no_grad self-cond signal) ALWAYS runs in the forward. Always running pass-1 makes the FSDP collectives identical every step regardless of the coins (no rank desync) and keeps it correct for local_batch_size > 1 — a scalar per-microbatch coin only matched Google’s per-example mix when B == 1. Seeded by (step, microbatch_idx) for reproducibility; rank-correlation of the coins is harmless now that pass-1 no longer branches on them.

nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe._forward_backward_step(
idx,
batch,
loss_buffer,
num_diffusion_tokens,
num_ar_tokens = None,
num_batches,
is_train: bool = True
)

Response-window forward/backward: canvas-only logits + canvas-sliced loss.

nemo_automodel.recipes.dllm.train_ft.DiffusionGemmaSFTRecipe.setup()

Inject a dummy mask_token_id (unused by block_diffusion) then build.

class nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe()

Bases: TrainFinetuneRecipeForNextTokenPrediction

Recipe for dLLM (diffusion LLM) supervised fine-tuning.

Extends the standard fine-tuning recipe by:

  1. Wrapping the dataloader collate function to produce unshifted batches
  2. Applying token corruption before each forward pass
  3. Using dLLM-specific loss functions via a pluggable strategy
_response_window_collation
bool = False
nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._apply_corruption(
input_ids,
loss_mask,
microbatch_idx: int = 0
)

Apply token corruption via the configured strategy.

Parameters:

input_ids

Clean token IDs, shape [B, L].

loss_mask

Supervised positions mask, shape [B, L].

microbatch_idx
intDefaults to 0

Index of this microbatch within the step.

Returns:

Tuple of (noisy_input_ids, noise_mask, p_mask).

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._augment_batch_for_model(
batch,
clean_input_ids,
loss_mask
)

Add any model-specific forward inputs derived from the batch.

Default is a no-op; block_diffusion overrides this to attach the block-causal attention masks and canvas position ids.

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._compute_loss_denominators(
batches,
num_noise_tokens,
num_supervised_tokens
)

Return (num_diffusion_tokens, num_ar_tokens) — the GLOBAL token denominators for this step’s diffusion + AR losses.

Base: the pre-mask counts (diffusion per the strategy’s normalization_mode, AR = supervised). Subclasses whose final loss masks differ from these raw counts (e.g. block-diffusion, which restricts the diffusion loss to one selected canvas block and drops padding) override this to count the final masks — otherwise the losses are divided by tokens that cannot contribute and are silently under-scaled.

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._forward_backward_step(
idx,
batch,
loss_buffer,
num_diffusion_tokens,
num_ar_tokens = None,
num_batches,
is_train: bool = True
)

Override: apply dLLM corruption and compute dLLM loss.

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._run_train_optim_step(
batches,
max_grad_norm: float | None = None
)

Execute a single training step with dLLM loss.

Follows the parent pattern but uses loss_mask from the collate wrapper instead of labels != -100 for token counting.

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._run_validation_epoch(
val_dataloader
)

Run one validation pass with dLLM corruption and loss.

Computes per-batch loss with proper denominators, then accumulates weighted by noise token count to produce a per-noise-token average across the val set.

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe._wrap_dataloader_collate()

Replace dataloader collate functions with the dLLM single-pass collater.

Uses DLLMCollator which goes directly from variable-length sample lists to block-aligned tensors in one pass.

Requires datasets to produce unshifted format (input_ids + loss_mask, via _package_tokenized_example(unshifted=True)).

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe.log_train_metrics(
log_data
)

Log dLLM-specific training metrics.

nemo_automodel.recipes.dllm.train_ft.DiffusionLMSFTRecipe.setup()

Build all training components, then apply dLLM-specific overrides.

nemo_automodel.recipes.dllm.train_ft.corruption_sample_seed(
base_seed: int,
step: int,
dp_rank: int,
dp_size: int,
local_batch_size: int,
grad_acc_steps: int,
microbatch_idx: int,
offset: int
) -> int

Topology-independent per-sample corruption seed.

StatefulDistributedSampler shards the shuffled stream strided (rank r owns shuffled[r::dp_size]), so an example at (step, dp_rank, microbatch, offset) sits at global shuffled position step*gbs + dp_rank + (microbatch*local_batch_size + offset)*dp_size with gbs = local_batch_size * dp_size * grad_acc_steps. Seeding by that global index makes the noise an example receives independent of parallel topology (dp_size / local_batch_size / grad-accum) and resume-safe (a pure function of (step, sample), never the global RNG state).

Parameters:

base_seed
int

Run seed (self._self_cond_base_seed).

step
int

Optimizer step index.

dp_rank
int

Data-parallel rank (sampler convention; TP/CP peers share it).

dp_size
int

Number of data-parallel shards (sampler num_replicas).

local_batch_size
int

Per-rank micro-batch size.

grad_acc_steps
int

Gradient-accumulation micro-batches per optimizer step.

microbatch_idx
int

Micro-batch index within the step.

offset
int

Sample offset within the micro-batch.

Returns: int

A 63-bit non-negative seed for torch.Generator.manual_seed.

Topology-independence assumes the shuffled StatefulDistributedSampler (dataloader.group_by_length=false). With group_by_length=true the LengthGroupedSampler shards the length-sorted order, so the global-batch composition itself depends on dp_size and cross-topology reproducibility does not hold there. The two guarantees that do NOT depend on the sampler — TP/CP peers (which share dp_rank) drawing identical noise, and resume reproducing the same noise — hold regardless.

nemo_automodel.recipes.dllm.train_ft.main(
config_path = None
)

Main entry point for dLLM SFT recipe.

nemo_automodel.recipes.dllm.train_ft._CORRUPTION_SEED_MULT = 2654435761
nemo_automodel.recipes.dllm.train_ft._CORRUPTION_STREAM_OFFSET = 2 << 42
nemo_automodel.recipes.dllm.train_ft._SEED_MASK = (1 << 63) - 1
nemo_automodel.recipes.dllm.train_ft.logger = logging.getLogger(__name__)