Masked diffusion models (MDMs) generate text by unmasking several tokens per step, but they are trained and sampled under different conditions. The model is trained on randomly masked sequences, whereas inference follows a trajectory shaped by the model's own predictions. Additionally, each step has no access to what the previous one computed. Recent methods narrow these limitations from separate angles, leaving open how these choices interact. We introduce PUMBA, a unified framework for trajectory-aware training that trains the denoiser on consecutive steps of policy-induced trajectories, passes information between steps, and optimizes them jointly by backpropagation through time. A controlled study of this design space shows that i) exact train--inference alignment fails due to local overfitting, whereas a looser alignment still brings training masks closer to those seen at inference; ii) passing continuous information outperforms discrete gradient estimators through the commitment at each step; and iii) performance improves as backpropagation through time spans more steps, which we support theoretically. Combined, these components match the best checkpoint of a same-size autoregressive model. Building on these findings, we scale PUMBA to supervised fine-tuning of LLaDA-8B, where it improves the trade-off between performance and number of function evaluations (NFEs) in both full-canvas and block diffusion generation. At matched performance, it needs up to 22% fewer NFEs than standard fine-tuning with twice the budget in full-canvas generation, and up to 26% fewer than standard fine-tuning for the same number of steps in block diffusion.
Figures & tables
PU + REINFORCE
41.2
PU + Gumbel-softmax
42.4
PU + Straight-through
44.7
PU + Carry
51.2
Table 1 : Peak GSM8K accuracy (%) for different BPTT techniques at W=2 on TinyGSM.
AR
55.3
MDM
34.8±2.5
MDM ×2 epochs
42.6
MDM ×8 epochs
43.6
PU
40.2±0.3
PU + Carry, W=1
44.4
PU + Carry, W=2
48.7
Table 2 : Best GSM8K accuracy (%) at u=2 ; MDM and PU: mean ± s.d. over three seeds.
Appendix figures & tables11 assets
Supplementary material from the paper’s appendix.
Appendix
TinyGSM, Kim et al. (2026)
LLaDA-8B, Section 4
Start of a trajectory
A random fraction r∼U[0,1/K) of the L positions revealed, chosen uniformly at random
Fully masked
Progress index
Stage p∈{0,…,K−1}
None
Positions committed per step
Enough to reach round(rL) revealed positions, r∼U[(p+1)/K,(p+2)/K) , capped at L−1
The u most confident
Confidence threshold
Every other position above τ=0.9 ; the stage is then recomputed from the revealed fraction
Every other position above τ=0.9
Schedule
K raised stepwise from 12 to 42 ( Table 4 )
u fixed per run ( Table 7 )
Retirement
After the step taken at stage K−1
Once the response and its EOS are fully revealed
Appendix
Table 3 : The two progressive unmasking constructions. L is the number of maskable positions of a sample and K the number of stages. A trajectory is at stage p when a fraction in [p/K,(p+1)/K) of its L positions is revealed.
Final hidden state; zero-initialized LayerNorm; zero at the first step; fp32 storage; p=1
Appendix
Table 7 : Post-SFT stage of LLaDA-8B. Wall-clock time is the training run alone, without evaluation, on 64 H100 unless stated.
Full canvas
Block diffusion
Content (prompt + response + EOS)
1.81 B (2.71 B)
1.81 B (10.84 B)
of which prompt
1.22 B (1.83 B)
1.22 B (7.34 B)
of which response + EOS
0.58 B (0.88 B)
0.58 B (3.50 B)
EOS filler
—
0.03 B (0.20 B)
In the loss
0.58 B (0.88 B)
0.62 B (3.71 B)
Attended
2.32 B (3.49 B)
2.46 B (14.75 B)
Appendix
Table 8 : Token budget of LLaDA-8B SFT, per epoch over Dolci-Instruct-SFT, with the total of the reported run in parentheses: 1.5 epochs on the full canvas (25,195 updates of 128 sequences) and 6 epochs under block diffusion (100,781 updates). Counts are measured over one epoch of the training dataloader and, under block diffusion, of the vectorized layout the training step builds from it. Content is prompt plus response plus the terminal EOS after truncation to the 4096-token canvas, EOS filler completes the last block of each response, in the loss counts the positions the loss is computed on, and attended counts the positions the model attends over; under block diffusion, these are the context and the clean and noised copies of the response. Total canvas counts every sample as a full canvas, whatever part of it the sample uses: 4096 positions on the full canvas and 8192 under block diffusion, whose layout holds the clean and noised copies side by side.
Distinct
Relative completion loss
Validation
J
sequences
Minimum
At exit
300 later
completion loss
PU, u=2
231
3.6 k
0.54
0.88
0.98
6.92
PU, u=4
115
6.9 k
0.70
0.71
0.98
6.58
PU, u=8
57
13.7 k
0.75
0.75
0.98
5.72
PU, u=64
7
110 k
0.97
0.99
1.00
4.37
PU, u=128
4
192 k
0.98
1.00
1.03
4.36
Appendix
Table 9 : Runs of Figure 2 . Distinct sequences: training sequences visited over the 3,000 updates. Relative completion loss of the inserted sequences: its minimum, its value when the sequence leaves the batch ( J updates after insertion, one update for MDM), and its value 300 updates later. Validation completion loss: mean over the last 200 updates.
u=2
u=4
u=8
PU
6.92
6.58
5.72
PU schedule, MDM masks
7.61
6.67
6.23
Spread visits
9.70
7.18
6.03
Initial stages 0 to 31
9.71
7.64
5.73
No 1/t factor
7.08
6.44
5.81
Batch of 128 sequences
8.22
7.12
6.24
Appendix
Table 10 : Validation completion loss when one aspect of PU is changed at a time (MDM: 4.72 ).
Full canvas
u=16
u=32
u=64
u=128
PU
62.5(+1.6)
62.8(+1.9)
61.9(+1.0)
61.2(+0.3)
PU + Carry, W=1
61.9(+1.0)
62.6(+1.7)
62.4(+1.5)
61.7(+0.8)
PU + Carry, W=2
61.8(+0.9)
61.4(+0.5)
62.5(+1.6)
62.0(+1.1)
PU + Carry, W=4
61.9(+1.0)
63.0(+2.1)
62.6(+1.7)
62.6(+1.7)
PU + Carry, W=8
62.6(+1.7)
63.5(+2.6)
63.0(+2.1)
62.4(+1.5)
Appendix
Table 11 : Best score of each full-canvas Post-SFT variant, for every training u of the BPTT sweep, with the change relative to the departure SFT checkpoint in parentheses. Best and 2nd best highlighted.
Block diffusion
u=4
u=8
u=16
PU
62.8(+1.4)
61.9(+0.5)
62.1(+0.8)
PU + Carry, W=1
61.1(−0.2)
62.0(+0.6)
62.1(+0.7)
PU + Carry, W=2
63.8(+2.5)
62.0(+0.6)
63.6(+2.2)
PU + Carry, W=4
62.3(+0.9)
62.9(+1.5)
—
PU + Carry, W=8
62.1(+0.8)
—
—
Appendix
Table 12 : Best score of each block diffusion Post-SFT variant, for every training u of the BPTT sweep, with the change relative to the departure SFT checkpoint in parentheses. Best and 2nd best highlighted.
Model
Source
Shots
IFEval
GSM8K
MBPP
Published, reference only
LLaDA-8B-Base
LLaDA
n/r / 4 / 4
n/r
70.3
40.0
LLaDA-8B-Instruct, full canvas
LLaDA
n/r / 4 / 4
n/r
69.4
41.0
LLaDA-8B-Instruct, block 32
LLaDA
n/r / 4 / 4
n/r
77.5
34.2
LLaDA-8B-Base
Dream
n/r / 8 / 4
n/r
70.9
39.0
LLaDA-8B-Instruct
Dream
n/r
59.9
78.6
34.2
Appendix
Table 13 : LLaDA-8B: published figures against our own measurements. The published rows are reference only and are not protocol-matched to anything else in this paper. Shot counts are given per row as IFEval/GSM8K/MBPP, with n/r where the source does not report a value. Our rows decode with Fast-dLLM at τ=1 , left to right in blocks of 32 tokens, and report IFEval 0-shot prompt-level strict accuracy, GSM8K 8-shot exact match, and MBPP 3-shot pass@1. The two LLaDA-8B-Base rows differ only in the prompt format; the two SFT rows are the checkpoints the Post-SFT runs start from, evaluated in the chat format. The last two published rows are LLaDA as evaluated by Ye et al. (2025) , not Dream’s own model.
Masked diffusion models (MDMs) have emerged as a promising alternative to autoregressive models (ARMs) for language modeling. However, MDMs are known to learn substantially more slowly than ARMs, which may become problematic when scaling MDMs to larger models. Therefore, we ask the following question: how can we accelerate standard MDM training while maintaining its final performance? To this end, we first provide a detailed analysis of why MDM training is slow. We find that the main factor is the locality bias of language: the predictive information for a token is concentrated in nearby positions. We further investigate how this bias slows learning and suggest a simple yet effective remedy: bell-shaped time sampling as a training strategy. Notably, MDMs trained with our training recipe reach the same validation negative log-likelihood (NLL) up to ∼4× faster than standard training on One Billion Word Benchmark (LM1B). We also show faster improvements in generative perplexity, zero-shot perplexity, and downstream task performance on various benchmarks.
Masked Diffusion Models (MDMs) have emerged as a promising alternative to autoregressive models in language modeling, offering the advantages of parallel decoding and bidirectional context processing within a simple yet effective framework. Specifically, their explicit distinction between masked tokens and data underlies their simple framework and effective conditional generation. However, MDMs typically require many sampling iterations due to factorization errors stemming from simultaneous token updates. We observe that a theoretical lower bound of the factorization error exists, which standard MDMs cannot reduce due to their use of a deterministic single-state mask. In this paper, we propose the Infinite Mask Diffusion Model (IMDM), which introduces a stochastic infinite-state mask to mitigate the theoretical bound while directly inheriting the benefits of MDMs, including the compatibility with pre-trained weights. We empirically demonstrate that MDM fails to perform few-step generation even in a simple synthetic task due to the factorization error bound, whereas IMDM can find an efficient solution for the same task. Finally, when equipped with appropriate distillation methods, IMDM surpasses existing few-step distillation methods at small step counts on LM1B and OpenWebText. Code is available at https://Ugness.github.io/official_imdm.
Jaehoon Yoo, Wonjung Kim, Chanhyuk Lee +1
Korea Advanced Institute of Science and Technology (KAIST)
Masked diffusion language models (MDLMs) re-predict every position at each denoising step, but standard samplers commit tokens once revealed, leaving this revision capability unused. Existing approaches either add heuristic or learned mechanisms to revise committed tokens, or remask them back to [MASK] before re-predicting; a principled sampler that directly revises visible tokens without auxiliary modules remains underexplored. We introduce D3IM, a parameter-free sampler derived as a corrector-style reverse update that permits direct visible-to-visible revision without additional modules or auxiliary passes. D3IM also reveals a model-side obstacle we term preservation bias: the model tends to reproduce its own wrong committed tokens rather than correct them. We address this with SCOPE (Self-Conditioned On Prediction Errors), a lightweight post-training procedure that simulates D3IM's sampling process. On LLaDA-8B at 64 denoising steps, SCOPE+D3IM improves over the original LLaDA-8B with standard unmasking by +13.0 on GSM8K (68.3%), +4.8 on MATH-500 (23.6%), +15.3 on HumanEval (29.3%), and +10.4 on MBPP (30.8%), with gains that increase as more denoising steps are used on math and HumanEval.