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.