nemo_automodel.recipes.dllm.train_ft
nemo_automodel.recipes.dllm.train_ft
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
Functions
Data
API
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 position0is the first response token, so the diffusion block boundaries align to the response — blockiis 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_maskwith per-exampleprefix_lengths = prompt length,response_length = response length,enc_len = full sequence length;decoder_position_idsare 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_stepso 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 callsmodel(**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.
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.
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.
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.
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.
Response-window forward/backward: canvas-only logits + canvas-sliced loss.
Inject a dummy mask_token_id (unused by block_diffusion) then build.
Bases: TrainFinetuneRecipeForNextTokenPrediction
Recipe for dLLM (diffusion LLM) supervised fine-tuning.
Extends the standard fine-tuning recipe by:
- Wrapping the dataloader collate function to produce unshifted batches
- Applying token corruption before each forward pass
- Using dLLM-specific loss functions via a pluggable strategy
Apply token corruption via the configured strategy.
Parameters:
Clean token IDs, shape [B, L].
Supervised positions mask, shape [B, L].
Index of this microbatch within the step.
Returns:
Tuple of (noisy_input_ids, noise_mask, p_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.
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.
Override: apply dLLM corruption and compute dLLM loss.
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.
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.
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)).
Log dLLM-specific training metrics.
Build all training components, then apply dLLM-specific overrides.
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:
Run seed (self._self_cond_base_seed).
Optimizer step index.
Data-parallel rank (sampler convention; TP/CP peers share it).
Number of data-parallel shards (sampler num_replicas).
Per-rank micro-batch size.
Gradient-accumulation micro-batches per optimizer step.
Micro-batch index within the step.
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.
Main entry point for dLLM SFT recipe.