Native diffusion language-model training

Train supported diffusion checkpoints with diffusion_lm: and the usual Axolotl chat datasets. No plugin is required. The model profile selects its noise process, attention layout, logit alignment and reference loss.

Model Noise and layout Default native objective
Nemotron-Labs-Diffusion (3B, 8B) Absorbing mask, full sequence, aligned logits Inverse-time weighted CE at masked positions
Dream v0 Absorbing mask, full sequence, shifted logits Source-compatible masked CE with inverse-time weighting
DiffusionGemma Uniform token noise, causal encoder and bidirectional canvas Canvas CE, encoder AR CE and two-pass self-conditioning

Other models can still train with the legacy masked-diffusion objective by setting diffusion_lm.from_causal_lm: true with the legacy diffusion plugin; any other diffusion_lm: config on a model without a native profile is rejected rather than trained as causal SFT. See Converting a causal LM.

Installation

Native diffusion training runs on Axolotl’s supported Torch range, using the dependency versions pinned in pyproject.toml. attn_implementation: varlen additionally requires Torch 2.14, so install the torch==2.14.0 wheel for your accelerator from PyTorch’s installation guide before installing Axolotl:

python -m venv .venv-diffusion
source .venv-diffusion/bin/activate
python -m pip install --upgrade pip
python -m pip install torch==2.14.0  # or the accelerator-specific wheel
python -m pip install -e .

Put this environment’s bin directory first on PATH so axolotl train launches the matching Accelerate installation.

Start with Nemotron 3B

The small iteration target is nvidia/Nemotron-Labs-Diffusion-3B. From the repository root:

axolotl preprocess examples/nemotron-diffusion/lora-smoke.yaml
axolotl train examples/nemotron-diffusion/lora-smoke.yaml

This recipe trains two optimizer steps on 32 synthetic chat records. It checks preprocessing, packed training, gradient accumulation and adapter saving; it is not a quality benchmark. It pins the model and remote-code revision with revision_of_model and uses the published bidirectional diffusion mode, aligned logits and the existing mask token 100. Padding reuses the tokenizer EOS token without resizing the vocabulary. When you change models, remove or replace revision_of_model.

Attention backends

Native diffusion defaults to flex_attention when attn_implementation is omitted. eager and sdpa remain explicit fallbacks; the examples use sdpa as a dense reference. FlexAttention requires a compatible GPU, and its first execution includes kernel compilation.

Nemotron also supports PyTorch’s public variable-length attention:

attn_implementation: varlen
sample_packing: true

This path gathers valid Q/K/V tokens into THD format and passes cumulative document offsets to torch.nn.attention.varlen.varlen_attn, without building a dense document mask. PyTorch selects the underlying kernel. varlen requires Torch 2.14 or newer, a native full-sequence diffusion model (Nemotron), uncached bidirectional attention, and zero attention dropout. Dream and DiffusionGemma keep their existing backends: Gemma’s per-layer encoder-prefix/canvas visibility cannot be expressed as one full-sequence varlen call.

DiffusionGemma’s large attention heads can exceed a GPU’s per-kernel shared memory limit with the default Flex tile sizes. Smaller tiles can be configured without changing attention visibility:

attn_implementation: flex_attention
flex_attn_compile_kwargs:
  fwd_BLOCK_M: 16
  fwd_BLOCK_N: 16
  fwd_num_stages: 1
  bwd_BLOCK_M1: 16
  bwd_BLOCK_N1: 16
  bwd_BLOCK_M2: 16
  bwd_BLOCK_N2: 16
  bwd_num_stages: 1

Packed training

Enable sample_packing: true. Logical documents remain isolated even when they share a physical row, and position IDs restart at document boundaries. DiffusionGemma also keeps each canvas associated with its own encoder history.

sample_packing: true
sequence_len: 256
micro_batch_size: 2
diffusion_lm:
  overflow_policy: error
  generate_samples: false

sequence_len limits each example’s logical layout. For packed and batch_flattening rows, Axolotl derives the total allocation as sequence_len * micro_batch_size. Ordinary padded batches have no aggregate row budget. Evaluation uses eval_batch_size when it is set.

FlexAttention rounds capacity down to 128-token buckets. Encoder/canvas packing reserves another 128 tokens because the two streams pad independently. Use a derived allocation of at least 128 for full-sequence models and 256 for DiffusionGemma; increase sequence_len if individual examples do not fit. SDPA and eager use the full derived capacity.

The default overflow_policy: error raises on an oversized example. overflow_policy: drop filters and counts oversized examples during preprocessing; it never truncates them.

Dream can use eos_tail: visible_supervised to extend each logical example to sequence_len with supervised EOS tokens. These are semantic tokens and consume packing capacity. Physical padding is excluded from attention and loss.

Packing mainly saves memory; throughput depends on your lengths and attention backend, so benchmark both settings for your workload.

DiffusionGemma attention boundaries

Prompt prefill is causal. During denoising, each token in the current output canvas attends to the clean prefix and every valid token in that same noisy canvas. Bidirectional visibility applies to the output canvas, not the combined prompt and output sequence.

Training uses the same boundary. The encoder can process the complete clean sequence for its auxiliary autoregressive loss, but the decoder mask exposes only keys before the selected output block. Clean copies of that block and later output tokens stay hidden, and causal encoder attention prevents those future tokens from entering the visible prefix states.

Full-attention decoder layers see the entire eligible prefix. Sliding-attention layers see the retained prefix window and the entire current canvas; they do not apply a query-relative sliding window inside the canvas. Both layer types exclude other packed examples. Dream and Nemotron use their full-sequence layouts.

Objective and recurrence options

Unset options inherit the model’s reference behavior. Dream’s CART variant uses time_weighting: cart and cart_p; its optional focal reweighting uses token_reweighting: true, alpha and gamma.

DiffusionGemma’s default self-conditioning probability (self_conditioning.p) is 0.5. Its optional Rhine objectives use time_weighting: loo or inv_one_minus_t with objective_reduction: example_mean; rhine_weight_clip caps the 1 / (1 - t) weights. encoder_ar_weight scales the separate encoder loss. self_conditioning.train_module: true includes the self-conditioning module in the saved adapter.

diffusion_lm:
  unroll:
    k_max: 2
    grad_through_steps: false

With unroll.k_max above 1, each step draws a read count uniformly from 1 to k_max. Earlier reads run without retained gradients and the final read supplies the loss. grad_through_steps: true is available only for models with self-conditioning.

Adapters, merging and inference

Native training requires LoRA with explicit text attention or dense-MLP projection targets. DoRA, layer replication, lora_target_linear and lora_target_parameters are rejected, and native vocabulary resizing is disabled. Nemotron additionally supports adapter: qlora (4-bit; 8-bit is rejected), FSDP2 via fsdp_config, and the fused LoRA kernels (lora_qkv_kernel, lora_o_kernel, lora_mlp_kernel) together with lora_fp32_gradients; Dream and DiffusionGemma remain unquantized LoRA without FSDP.

DiffusionGemma’s encoder and decoder text weights are tied, so their LoRA factors are shared. Its targets must be a single expression naming both the encoder language model and the decoder, for example:

lora_target_modules: '^model\.(encoder\.language_model|decoder)\.layers\.\d+\.self_attn\.(q|k|v|o)_proj$'

Vision, router and expert parameters stay frozen.

Nemotron supports opt-in Cut Cross Entropy for its aligned diffusion objective and for typed-decision training. With Axolotl’s CCE fork installed, add:

plugins:
  - axolotl.integrations.cut_cross_entropy.CutCrossEntropyPlugin
cut_cross_entropy: true

Keep any other plugins the recipe requires in that list. The final training read computes full-vocabulary CE from hidden states without materializing vocabulary logits; decision training still computes its restricted-label and Brier terms over the allowed answers. Targets stay aligned, with no causal shift. This path requires a plain linear output head and does not support a LoRA output head or the loo objective.

axolotl merge-lora config.yaml exports a merged checkpoint. When merge_method is omitted, DiffusionGemma selects legacy so its tied factors merge once; setting merge_method: memory_efficient for DiffusionGemma raises.

For interactive generation, use axolotl inference config.yaml --chat. Each reply is a completion block appended to the conversation and denoised in one piece; see Inference. Native models generate completions only.

Typed-decision training

The diffusion_decision plugin trains a native diffusion model to answer typed multiple-choice, score and yes/no questions. Each record renders a fixed decision system prompt from its question schema, sends its state as the user message, and places one answer label per question in the output canvas using the canvas template of the public mmastrac/djev project.

examples/nemotron-diffusion/decision-lora-8b.yaml is the starting recipe for nvidia/Nemotron-Labs-Diffusion-8B. It uses BF16, attn_implementation: varlen, and rank-64/alpha-128 LoRA over all seven attention and MLP projections with loraplus_lr_ratio: 8. It trains one epoch with AdamW at 6.5e-6, microbatch 16 and eight accumulation steps (effective batch 128). The decision loss combines candidate-restricted CE, full-vocabulary CE and a 0.1 Brier term, with unroll.k_max: 2, no latent slots and non-reentrant gradient checkpointing. The YAML leaves the model revision unpinned; add revision_of_model to a local copy for an exact reproduction.

axolotl preprocess examples/nemotron-diffusion/decision-lora-8b.yaml
axolotl train examples/nemotron-diffusion/decision-lora-8b.yaml

Set the data_files of the datasets and test_datasets entries to separate normalized train and development JSONL files. The recipe evaluates every 50 steps and keeps the adapter with the best eval_loss. Keep held-out test data out of both entries.

Dataset format

type: diffusion_decision.jsonl reads already-normalized JSONL, one record per line. diffusion_decision.typed_decisions and diffusion_decision.procedural normalize rows from the typed_decisions and procedural typed-decision source schemas into the same form.

{"source":"support_demo","family":"returns","id":"case-0042","group":"case-0042","instructions":"Decide from the supplied case only.","state":{"customer":"A customer says their unopened order arrived yesterday and asks for a refund.","policy":"Unopened items delivered within 30 days may be refunded."},"questions":{"refund":{"type":"choice","instructions":"Choose the next action.","options":[{"name":"approve","description":"Approve a refund under the stated policy."},{"name":"deny","description":"Deny the request because the policy does not apply."},{"name":"escalate","description":"Send the case to a human reviewer."}]},"urgency":{"type":"score","instructions":"Rate how urgently this needs review.","levels":["routine","soon","urgent"]},"eligible":{"type":"noul","instructions":"Is the customer eligible for a refund?","criteria":{"true":"The stated policy permits a refund.","false":"The stated policy does not permit a refund."}}},"labels":{"refund":{"kind":"dist","probs":[0.78,0.02,0.20]},"urgency":{"kind":"hard","gold_idx":0},"eligible":{"kind":"set","allowed_set":[0]}},"source_metadata":{"provenance":{"dataset":"my-org/support-decisions","revision":"immutable-revision","split":"train","source_id":"case-0042"}}}

source, id and group identify the row; family is optional and is used for split isolation. Records with the same source, family, state and top-level instructions that share a group are rendered on one canvas, up to max_questions_per_canvas questions. There is no per-row system field: put shared context in state, global instructions in instructions, and per-question instructions on each question. Keep question IDs stable and nonempty.

Type Required schema field Canonical alternative order
choice options List order. An option is a string or an object with name and description; descriptions are rendered.
score levels List order, lowest to highest as defined by the source. Do not sort levels lexically.
noul none yes, then no. Optional criteria documents the true/false conditions and is rendered.

Label indexes always refer to this canonical order, never to option names or token IDs. Answer tokens come from labels.codebook: vendored26 uses the djev letters A-Z (up to 26 alternatives); expanded52 allows up to 52.

Labels

Every question has exactly one label:

{"kind":"hard","gold_idx":1}
{"kind":"dist","probs":[0.10,0.70,0.20]}
{"kind":"set","allowed_set":[0,2]}

hard names one zero-based alternative. set is a nonempty set of unique valid indexes. dist has one finite, nonnegative probability per alternative that sums to one within 1e-4; zero entries are valid and must not be omitted. Partial distributions, ranking targets, per-row candidate token IDs and hard/soft blending are not accepted.

labels.label_softmax selects candidate-restricted CE (restricted), full-vocabulary CE over the allowed label tokens (full), or both (the default). The full-vocabulary term still penalizes probability outside the valid answers. full_ce_weighting: dft applies DFT weighting to hard-label full CE. hard_label_smoothing spreads that fraction of a one-hot CE target uniformly across the question’s valid answer tokens; Brier targets are unchanged.

Splits and provenance

Keep train, development and test data in separate files or dataset entries, and use split: train only for training entries. Preserve source_metadata.provenance with an immutable source location and revision, the split, and the source row ID. typed_decisions rows train only from a declared train split and reject a row split that disagrees with it. After normalization, the loader rejects states or grouped contexts shared between splits.

Mixture and weighting

mixture.weights sets per-source sampling weights by source. Without weights, each source is sampled in proportion to its row count raised to mixture.temperature (default 0.5). max_examples_per_source caps rows per source. loss_weight multiplies each source’s loss. Loss is averaged over the questions on a canvas and then over canvases, so a canvas with many questions does not outweigh a single-question canvas.

To train on a prebuilt draw sequence, set mixture.premixed: true and per_batch_stratified: false. Each training row then needs a unique draw ID and an origin matching its own identity; copies of one origin must be identical apart from this metadata:

"source_metadata": {
  "premix": {
    "draw_id": "run-000001",
    "origin": {"source": "support_demo", "id": "case-0042", "group": "case-0042"}
  }
}

Batching

batch_flattening: true concatenates the logical microbatch and keeps decisions isolated through the varlen attention path. For variable-length decisions, set batch_flattening: false, sample_packing: true and eval_sample_packing: false instead. Packing uses sequence_len as the per-decision limit and sequence_len * micro_batch_size as the physical token budget. It requires mixture.per_batch_stratified: false and does not support sampled latent-slot counts. Reduce micro_batch_size and raise gradient_accumulation_steps if the effective batch does not fit.

Outputs

Each saved adapter includes diffusion_decision_manifest.json, which records the base model and revision, diffusion spec, canvas width, label codebook, latent slots, read protocol and tokenizer special IDs that an external inference integration must reproduce. Axolotl does not ship a decision-serving endpoint. Use a separate config for a frozen test split, and do not use test data for checkpoint selection.

Public procedural decision mix

scripts/diffusion_lm/build_public_decision_mix.py materializes a reproducible train/dev mix from tasksource/procedural-typed-decisions at revision 916e6cce365a65c37d58db70c9e369795817651e:

python scripts/diffusion_lm/build_public_decision_mix.py ./data/public-procedural
axolotl train examples/nemotron-diffusion/decision-lora-8b-public-procedural.yaml

The output contains only rows from that dataset plus a contract.json recording the pinned revision and licenses. The dataset card declares Apache-2.0; the separate generator repository publishes its code under CC-BY-4.0. Keep contract.json with any materialized output you share. The model is covered by its own model-card license.

--protected-normalized PATH (repeatable) excludes any public row whose canonical state matches a row in your normalized evaluation JSONL; that file is never copied to the output. This is canonical-state decontamination only, not cross-source group or family disjointness. The train and dev outputs are checked against each other for state, family and group overlap.

decision-lora-8b-public-procedural.yaml pins the model revision and sets mixture.temperature: 1.0, which schedules each usable row once. At the pinned revisions, preprocessing keeps 30,080 train rows at the 2,048-token budget, which is 235 effective batches of 128, hence max_steps: 235. This mix is a reproducible starting point for the format, not a quality result.

Converting a causal LM

diffusion_lm.from_causal_lm: true trains an ordinary causal checkpoint with the legacy masked-diffusion objective, controlled by noise_schedule, min_mask_ratio, max_mask_ratio, importance_weighting and the mask-token options. It requires plugins: [axolotl.integrations.diffusion.DiffusionPlugin], which supplies the diffusion trainer; without it the config is rejected. Existing configurations with a deprecated diffusion: block still load through the same plugin, which translates them to this form.