# On Trajectory-Aware Training for Masked Diffusion Language Models

> PUMBA trains masked diffusion language models on inference-like trajectories via continuous hidden-state carries and backpropagation through time, matching autoregressive accuracy while decoding multiple tokens per step.

- **Source:** [arXiv](https://arxiv.org/abs/2609.37974)
- **Published:** 2026-10-03
- **Permalink:** https://picx.dev/p/j3HP5f
- **Whiteboard:** https://picx.dev/p/j3HP5f/image

## Summary

## 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** $W$ 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 $u$) 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 $W$, 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:

1. **Distribution shift**: The model and policy jointly shape inference trajectories whose states differ from the training distribution of random masks.
2. **No memory across steps**: Since training sequences are drawn independently, the model never learns to use information from previous steps.

**Notation.** Let $\mathcal{V}$ be a finite vocabulary containing a mask token $m \in \mathcal{V}$, and $x = (x_1, \ldots, x_L) \in \mathcal{V}^L$ a sequence of length $L$ with masked positions $\mathcal{M}(x) = \{i : x_i = m\}$. We write $\Delta^d$ for the probability simplex in $\mathbb{R}^d$ and $\mathrm{sg}[\cdot]$ for stop-gradient.

**MDM training objective.** Under the linear noising schedule, each token of $x_0$ is replaced by $m$ independently with probability $t$ (Sahoo et al., 2024; Shi et al., 2024). The denoiser $f_\theta: \mathcal{V}^L \to (\Delta^{|\mathcal{V}|})^L$ returns a distribution $f^i_\theta(\cdot | x_t)$ over the clean token at each position $i$. Training maximizes an evidence lower bound reducing to weighted cross-entropy over masked positions:

$$\mathcal{L}_{\text{MDM}} = \mathbb{E}_{x_0, t, x_t} \left[ \ell(x_0, x_t) \right], \quad \ell(x_0, x_t) = \frac{1}{t} \sum_{i \in \mathcal{M}(x_t)} -\log f^i_\theta(x^i_0 | x_t) \tag{2.1}$$

The factor $1/t$ follows from the bound and normalizes the sum over the $L \cdot t$ positions masked in expectation.

**Inference.** Generation starts from the fully masked sequence and reveals tokens over multiple steps. Each step applies an unmasking policy $g$ that selects $S = g(x_{t_j}) \subseteq \mathcal{M}(x_{t_j})$ and fills each $i \in S$ with a token drawn from $f^i_\theta(\cdot | x_{t_j})$. The **top-$u$ policy** (Chang et al., 2022) reveals the $u$ masked positions with highest confidence $c_i = \max_{v \in \mathcal{V}} f^i_\theta(v | x_{t_j})$. For larger models, **Fast-dLLM** (Wu et al., 2026b) reveals every position with $c_i \geq \tau$, 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 $1/t$ factor with $t$ being the realized masking ratio.

### Carry (Continuous Information Passing)

The denoiser also takes the previous hidden state as input and returns its own:

$$f_\theta(\cdot | x_{t_j}, h_j), \quad h_{j+1} = f_{\text{carry}}(x_{t_j}, h_j) \tag{3.1}$$

The carry $h_j \in \mathbb{R}^{L \times d}$ 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 $x_0$.

### Backpropagation Through Time (BPTT)

When each update trains on a single step ($W = 1$), the model learns to use the carry it receives but not to produce a useful one. Computing the loss on $W$ consecutive steps and backpropagating through the carries between them optimizes the carry for subsequent predictions. Trajectories longer than $W$ are split into consecutive windows, each with its own update (truncated BPTT).

### Mask Discrepancy Metric

**Definition 1 (Mask discrepancy).** For a sample $x_0$ and masking ratio $t$, let $P_{\text{train}}$ and $P_{\text{inf}}$ be the distributions of masks $m \in \{0,1\}^L$ at ratio $t$ produced by training and sampling. The mask discrepancy is:

$$D_{\text{mask}}(t) = \mathbb{E}_{x_0} \left[ \text{MMD}^2_k(P_{\text{train}}, P_{\text{inf}}) \right], \quad k(m, m') = \exp\left(-d_H(m, m')/\sigma\right)$$

where MMD is maximum mean discrepancy (Gretton et al., 2012) and $d_H$ the Hamming distance. The kernel is strictly positive definite on $\{0,1\}^L$, so $D_{\text{mask}}(t) = 0$ exactly when training and inference produce the same masks at ratio $t$.

## Empirical Validation / Results

### Local Overfitting at Small $u$

At small $u$ (the targeted inference regime), PU breaks two assumptions of stochastic optimization: each sequence is reused $J$ times instead of once, and consecutive batches share all but a fraction $1/J$ of their sequences. This causes **local overfitting**: at $u = 2$, 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 $u$ than used at inference (the PUMA recipe of Kim et al., 2026) keeps alignment benefits without local overfitting. The effective training $u$ is ~133 at start and ~67 at end, versus $u = 2$ at inference. Despite this mismatch, PUMA still reduces $D_{\text{mask}}$ at $u = 2$ well below MDM (Figure 3).

### Comparison of Gradient Channels

With window fixed at $W = 2$ 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 $u=2$ |
|---|---|
| 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 |
| PU + Carry, $W=4$ | 52.6 |
| **PU + Carry, $W=8$** | **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 $W=2$ and $W=8$. 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 $W=1$), a step's loss cannot reward earlier steps; with BPTT, it reaches at most $W-1$ earlier steps.
(ii) Assume the carry is contractive with factor $\gamma < 1$. Let $\theta_W$ be a point where the truncated gradient vanishes and the loss satisfies a Polyak–Łojasiewicz condition. Then:

$$\text{KL}(p^\star \| p_{\theta_W}) - \inf_\theta \text{KL}(p^\star \| p_\theta) \leq C \left( \frac{J}{W - 1} \right)^2$$

where $J$ is the number of decoding steps and $C$ does not depend on $W$. The gap to the best achievable sampler shrinks as the window grows.

### Large-Scale SFT on LLaDA-8B

**Full-canvas generation**: PUMBA with $u=32$, $W=8$ 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 ($u=16$, $W=2$) 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 $u=1$, a run sees only 0.6% of samples visited by MDM, rising to 26% at $u=128$. Block diffusion is less exposed, reaching 100% by $u=32$ due to parallel-block stratification. As $u$ decreases, the model overfits locally and training collapses; at $u=128$, 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 $W$ grows, with a $\left(\frac{J}{W-1}\right)^2$ rate.
- Explains why continuous carries outperform discrete gradient estimators: no gradient flows through committed tokens without a carry, and BPTT reaches at most $W-1$ 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 $u$ 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:

1. **Alignment must be loose, not exact**: Local overfitting prevents training at inference-time reveal counts; training at larger $u$ still reduces mask discrepancy.
2. **Continuous carries are essential**: They outperform discrete gradient estimators by providing a differentiable channel through the discrete commitment.
3. **Larger BPTT windows help**: Performance improves with $W$, 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.

---

_Markdown view of https://picx.dev/p/j3HP5f, served by PicX — AI-generated visual whiteboard summaries of research papers._
