Full text not available for this paper
Summary (Overview)
- PUMBA (Progressive UnMasking with Backpropagation Across steps) is a unified framework for trajectory-aware training of masked diffusion language models (MDMs), addressing the train–inference gap by training on consecutive steps of policy-induced trajectories, passing continuous information between steps, and optimizing jointly via backpropagation through time (BPTT).
- The framework spans three design axes: the policy that builds training trajectories, the carry (continuous information passed between steps), and the window of consecutive steps optimized jointly.
- Key findings: (i) exact train–inference alignment fails due to local overfitting, but looser alignment (training at larger reveal counts ) still reduces mask discrepancy; (ii) continuous carries outperform discrete gradient estimators (REINFORCE, Gumbel-softmax, straight-through) for crossing the discrete commitment at each step; (iii) performance improves with larger BPTT windows , supported by a theoretical bound.
- PUMBA matches the performance of a same-size autoregressive model on TinyGSM/GSM8K while decoding more than one token per step.
- Scaled to supervised fine-tuning (SFT) of LLaDA-8B, PUMBA improves the performance–NFE (number of function evaluations) Pareto frontier: up to 22% fewer NFEs than standard SFT with twice the budget in full-canvas generation, and 26% fewer in block diffusion.
Introduction and Theoretical Foundation
Masked diffusion models (MDMs) generate text by iteratively unmasking several tokens per step. The training procedure, however, masks sequences randomly and independently, while inference follows a trajectory shaped by the model's own predictions via an unmasking policy. This creates two gaps:
- Distribution shift: The model and policy jointly shape inference trajectories whose states differ from the training distribution of random masks.
- No memory across steps: Since training sequences are drawn independently, the model never learns to use information from previous steps.
Notation. Let be a finite vocabulary containing a mask token , and a sequence of length with masked positions . We write for the probability simplex in and for stop-gradient.
MDM training objective. Under the linear noising schedule, each token of is replaced by independently with probability (Sahoo et al., 2024; Shi et al., 2024). The denoiser returns a distribution over the clean token at each position . Training maximizes an evidence lower bound reducing to weighted cross-entropy over masked positions:
The factor follows from the bound and normalizes the sum over the positions masked in expectation.
Inference. Generation starts from the fully masked sequence and reveals tokens over multiple steps. Each step applies an unmasking policy that selects and fills each with a token drawn from . The top- policy (Chang et al., 2022) reveals the masked positions with highest confidence . For larger models, Fast-dLLM (Wu et al., 2026b) reveals every position with , or the single most confident one if none qualifies.
Methodology
Progressive Unmasking (PU)
PU trains on the trajectory that the inference policy itself produces. For each sample, training starts from the fully masked sequence; at each optimizer step, the loss is computed on the current sequence, then the inference-time unmasking policy reveals more positions. This continues until the sample is fully revealed. Key characteristics:
- Teacher forcing: Revealed positions commit the ground-truth token (not the model's prediction), keeping targets consistent with visible context.
- Loss: Uses the MDM loss of Eq. (2.1) unchanged, including the factor with being the realized masking ratio.
Carry (Continuous Information Passing)
The denoiser also takes the previous hidden state as input and returns its own:
The carry is the last hidden state before the output projection. Each trajectory starts from a zero carry, and each inference step passes its carry to the next. The carry provides a weak form of student forcing: it reflects the model's imperfect predictions while visible tokens and targets stay consistent with .
Backpropagation Through Time (BPTT)
When each update trains on a single step (), the model learns to use the carry it receives but not to produce a useful one. Computing the loss on consecutive steps and backpropagating through the carries between them optimizes the carry for subsequent predictions. Trajectories longer than are split into consecutive windows, each with its own update (truncated BPTT).
Mask Discrepancy Metric
Definition 1 (Mask discrepancy). For a sample and masking ratio , let and be the distributions of masks at ratio produced by training and sampling. The mask discrepancy is:
where MMD is maximum mean discrepancy (Gretton et al., 2012) and the Hamming distance. The kernel is strictly positive definite on , so exactly when training and inference produce the same masks at ratio .
Empirical Validation / Results
Local Overfitting at Small
At small (the targeted inference regime), PU breaks two assumptions of stochastic optimization: each sequence is reused times instead of once, and consecutive batches share all but a fraction of their sequences. This causes local overfitting: at , the completion loss of a sequence falls sharply within ~50 steps of entering the batch, then rises again before leaving—the model fits the sequence on its trajectory's masks, after which the sequence provides no gradient for revealed tokens.
Mitigation: Training at a larger than used at inference (the PUMA recipe of Kim et al., 2026) keeps alignment benefits without local overfitting. The effective training is ~133 at start and ~67 at end, versus at inference. Despite this mismatch, PUMA still reduces at well below MDM (Figure 3).
Comparison of Gradient Channels
With window fixed at on TinyGSM:
| Technique | Peak GSM8K Accuracy (%) |
|---|---|
| PU + REINFORCE | 41.2 |
| PU + Gumbel-softmax | 42.4 |
| PU + Straight-through | 44.7 |
| PU + Carry | 51.2 |
The continuous carry outperforms all discrete gradient estimators.
Composing Design Choices
| Method | Best GSM8K Accuracy (%) at |
|---|---|
| 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, | 44.4 |
| PU + Carry, | 48.7 |
| PU + Carry, | 52.6 |
| PU + Carry, | 55.6 |
The three design choices compose: each adds accuracy on top of the previous ones. At matched denoiser passes, MDM trained 2× or 8× longer remains substantially below the corresponding BPTT at and . Together, the three choices match the best AR checkpoint while revealing more than one token per step.
Theoretical Support
Proposition 1 (Informal). Assume teacher-forced commits and a reveal rule matching inference. (i) Without a carry (or ), a step's loss cannot reward earlier steps; with BPTT, it reaches at most earlier steps. (ii) Assume the carry is contractive with factor . Let be a point where the truncated gradient vanishes and the loss satisfies a Polyak–Łojasiewicz condition. Then:
where is the number of decoding steps and does not depend on . The gap to the best achievable sampler shrinks as the window grows.
Large-Scale SFT on LLaDA-8B
Full-canvas generation: PUMBA with , gains 2.6 points in score over the 1.5-epoch SFT checkpoint—over three times the gain of doubling the SFT budget—and at matched score needs up to 22% fewer NFEs than the 3-epoch SFT.
Block diffusion: The best configuration (, ) gains 2.2 points in 4000 steps, more than the 50k SFT steps from 3 to 6 epochs, and at matched score needs 26% fewer NFEs than the Post-SFT MDM control.
Sample diversity trade-off: At full-canvas , a run sees only 0.6% of samples visited by MDM, rising to 26% at . Block diffusion is less exposed, reaching 100% by due to parallel-block stratification. As decreases, the model overfits locally and training collapses; at , PU only matches the Post-SFT MDM control.
Theoretical and Practical Implications
Theoretical contributions:
- Identifies local overfitting as the mechanism preventing exact train–inference alignment in PU, explaining why training at larger reveal counts is necessary.
- Provides a theoretical bound (Proposition 1) showing that the KL divergence to the optimal sampler shrinks as the BPTT window grows, with a rate.
- Explains why continuous carries outperform discrete gradient estimators: no gradient flows through committed tokens without a carry, and BPTT reaches at most earlier steps.
Practical implications:
- PUMBA provides a practical Post-SFT recipe: a short extra training stage (4000 steps) pushes a model beyond what additional SFT allows, and the resulting reduction in denoising steps amortizes the extra training cost.
- The framework is applicable to both full-canvas and block diffusion formulations, making it broadly relevant for modern MDM architectures.
- The trade-off between trajectory alignment and sample diversity is a crucial design consideration: training must be small enough to align with inference but large enough to allow effective learning.
Conclusion
PUMBA improves masked diffusion training by aligning it with the trajectories encountered during generation. The method caches the denoiser's hidden state across steps and uses truncated BPTT to train each step for its contribution to subsequent predictions. Key findings:
- Alignment must be loose, not exact: Local overfitting prevents training at inference-time reveal counts; training at larger still reduces mask discrepancy.
- Continuous carries are essential: They outperform discrete gradient estimators by providing a differentiable channel through the discrete commitment.
- Larger BPTT windows help: Performance improves with , supported by theory.
Limitations and future work:
- Local overfitting at small reveal counts remains a challenge; turning these findings into procedures that sustain closer inference alignment across model scales is open.
- Extending the carry's weak student forcing to sampled tokens did not improve on teacher forcing (irreversible wrong commitments leave context inconsistent with targets); remasking could address this.
- Richer carries (e.g., fixed-size working memory like MetaState), BPTT-specific gradient checkpointing, selective gradient propagation, or adaptive windows could improve the trade-off between training cost, memory, and credit-assignment horizon.
Related papers
- Identical Runs, Different Results: Benchmarking AI Coding Agents on Open-Weight Models
Identical runs of the same AI coding agent vary more than differences between agent-model pairings, so best-of-three with compliance checking beats single-run benchmarking.
- Cost-free Spectral Estimation for Adaptive Newton--Schulz in Matrix Optimizers
Newton–Schulz iterations expose spectral moments via free scalar reductions, enabling spectrum-adaptive routine selection that cuts polar error up to 90x and improves LLM pretraining loss.
- Audit the Scaffold, Not the Checkpoint: A Stationarity Dichotomy for Recursive Self-Improvement in Agentic Coding
Frozen weights do not guarantee safe saturation in agentic coding; auditors must monitor scaffold expansion, not just checkpoint freezing.