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.yamlThis 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: trueThis 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: 1Packed 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: falsesequence_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: falseWith 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: trueKeep 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.yamlSet 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.yamlThe 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.