← All Posts
Deep Learning · Diffusion & Flow Models · Large Language Diffusion Models· Part 2 of 8

The Masked Diffusion Objective, Derived

The $1/t$ weight has a probabilistic job. It comes from the rate at which a masked token must be revealed when reversing a linear masking process. This chapter derives it, then shows how to estimate the loss without accidentally changing the objective.

Fix the objects before differentiating

Use the absorbing path from Part 1: clean at $t=0$, fully masked at $t=1$, survival $a(t)$. Write $\rho(t)=1-a(t)$ for the masking probability. The denoiser $\mu_\theta^i(v\mid x_t,t)$ predicts clean vocabulary tokens, never $m$. A visible token is carried through unchanged. These constraints make several reverse-process terms vanish.

First solve the reverse problem with the answer supplied

Take two times $s<t$. If $x_t^i$ is visible, the earlier state is the same token. If $x_t^i=m$, two histories are possible: it was already masked at $s$, or it became masked between $s$ and $t$. Their unconditional probabilities are $1-a(s)$ and $a(s)-a(t)$. Divide by the probability $1-a(t)$ of being masked now:

$$q(x_s^i\mid x_t^i=m,x_0^i)=\frac{1-a(s)}{1-a(t)}\delta_m+\frac{a(s)-a(t)}{1-a(t)}\delta_{x_0^i}.$$

Define the reveal probability $r(s,t)=[a(s)-a(t)]/[1-a(t)]$. The two coefficients are nonnegative and sum to one. Replace the unavailable clean token with its learned posterior to obtain a normalized model transition:

$$p_\theta(x_s^i\mid x_t)=\begin{cases}\delta_{x_t^i},&x_t^i\ne m,\\ {}[1-r(s,t)]\delta_m+r(s,t)\mu_\theta^i(\cdot\mid x_t,t),&x_t^i=m.\end{cases}$$

For a linear schedule, $r(s,t)=(t-s)/t$. Moving from $t=0.8$ to $s=0.6$ reveals each current mask with probability $0.25$. It does not reveal 25% of the entire sequence. In expectation, the mask fraction changes from $0.8$ to $0.8(0.75)=0.6$.

Where cross-entropy enters the likelihood bound

For a finite grid, define a generative chain from the all-mask prior to clean text. Jensen's inequality applied to the latent corruption trajectory gives

$$-\log p_\theta(x_0)\le\mathbb E_{q(x_{t_1},\ldots,x_{t_K}\mid x_0)}\left[\log\frac{q(x_{t_1},\ldots,x_{t_K}\mid x_0)}{p_\theta(x_0,x_{t_1},\ldots,x_{t_K})}\right].$$

Rearranging this expression produces a terminal prior term, reverse-transition KL terms, and a reconstruction term. The all-mask prior matches the exact terminal forward distribution, so its KL is zero. At an intermediate masked coordinate, the true and modeled reverse transitions assign the same probability $1-r$ to staying masked. On a reveal, the true conditional transition selects $x_0^i$ and the model uses $\mu_\theta^i$. The coordinate KL is therefore

$$\mathrm{KL}\!\left((1-r)\delta_m+r\delta_{x_0^i}\;\middle\|\;(1-r)\delta_m+r\mu_\theta^i\right)=-r\log\mu_\theta^i(x_0^i\mid x_t,t).$$

The cancellation is exact because the clean vocabulary excludes $m$. Visible positions contribute zero. This is why letting the output head put probability on the mask symbol breaks the simplified expression.

Set $s=t-h$. For small $h$, $a(t-h)-a(t)=-a'(t)h+o(h)$. Summing KL contributions and taking the continuous-time limit gives the denoising term

$$\mathcal L(\theta)=\mathbb E_{x_0}\int_0^1\frac{-a'(t)}{1-a(t)}\;\mathbb E_{x_t\mid x_0}\left[\sum_{i=1}^{L}\mathbf1[x_t^i=m]\big[-\log\mu_\theta^i(x_0^i\mid x_t,t)\big]\right]dt.$$

Under exact absorbing endpoints, the carry-through parameterization, and a vanishing clean-end reconstruction contribution, this is the negative-ELBO objective for the continuous-time model. If an implementation truncates time or leaves nonzero endpoint noise, it must handle the corresponding boundary terms; it cannot label an arbitrary weighted loss “exact NLL.” The simplified formulations in MDLM and MD4 provide the formal background.

The linear schedule and its apparent singularity

With $a(t)=1-t$ and $t\sim\mathrm{Unif}(0,1)$, the weight is $1/t$. Dividing by a fixed sequence length gives a per-token training loss:

$$\mathcal L_{\mathrm{token}}=\mathbb E\left[\frac{1}{L t}\sum_i\mathbf1[x_t^i=m]\,\mathrm{CE}_i\right].$$

A given position is selected with probability $t$, so the indicator cancels the inverse masking probability in expectation. That does not mean the estimator is numerically harmless. In a toy case where every selected loss equals a constant $c$, its second moment contains $c^2/t$ and the integral diverges at zero. Finite mean and finite variance are different questions.

A small cutoff $t\ge\varepsilon$ is a practical approximation. If sampling uniformly on $[\varepsilon,1]$, multiply the average by $1-\varepsilon$ to estimate the truncated integral. State what was truncated; clipping a denominator is not a derivation of the full likelihood bound.

A useful alternative: force the scored position to be masked

For each position, condition on it being masked:

$$\mathbb E_{x_t\mid x_0}\left[\frac{\mathbf1[x_t^i=m]}{t}\,\mathrm{CE}_i\right]=\mathbb E_{x_t\mid x_0,\,x_t^i=m}[\mathrm{CE}_i].$$

Thus we can choose a uniformly random target position $I$, force that position to be masked, and corrupt every other position independently at rate $t$. Score only $I$. Averaged over $I$, this estimates exactly the same per-token objective for the linear schedule without an explicit $1/t$ multiplier. It spends one sequence forward pass on one supervised position, so it trades numerical simplicity for fewer labels per pass.

# x0: [batch, length], containing only clean vocabulary IDs
# model(xt): [batch, length, vocab]; mask is an INPUT-only ID
batch, length = x0.shape
t = torch.rand(batch, 1, device=x0.device)
target = torch.randint(length, (batch,), device=x0.device)
rows = torch.arange(batch, device=x0.device)
masked = torch.rand(batch, length, device=x0.device) < t
masked[rows, target] = True
xt = x0.masked_fill(masked, mask_id)
logits = model(xt)
loss = F.cross_entropy(logits[rows, target], x0[rows, target])

This code assumes all positions are valid, fixed-length data. With padding or prompt/response data, sample $I$ uniformly from the intended target positions and match the corruption distribution on the other editable positions. Part 6 defines that conditional task.

Why dividing by the realized mask count changes things

Suppose $M=\sum_i\mathbf1[x_t^i=m]$. Replacing the loss above with $\sum_i\mathbf1[x_t^i=m]\mathrm{CE}_i/M$ introduces a random denominator that depends on the corruption pattern. Hard examples with many masks and easy examples with few masks receive a different relative weighting. Skipping $M=0$ examples also changes the sampled time distribution unless corrected.

For a concrete calculation, let $L=1$, $t=0.2$, and the selected-token loss be $c$. The correct estimator $\mathbf1[\text{mask}]c/t$ has expectation $c$. Averaging only nonempty examples gives $c$ at this fixed $t$, but across a randomly sampled $t$, acceptance occurs with probability $t$: high-mask times are overrepresented. The apparent agreement in this tiny fixed-time example hides the bias of the complete procedure.

What the trained network estimates

At a fixed corrupted input, expected cross-entropy is minimized by the conditional distribution of its clean token. This is the categorical version of the conditional mean argument: the optimal prediction represents uncertainty over possible endpoints, rather than recovering the particular hidden answer every time.

Under a homogeneous token-independent masking schedule, the posterior over clean sequences given a particular mask pattern is independent of time. For every clean sequence compatible with the visible tokens, the factor $t^M(1-t)^{L-M}$ is the same and cancels in Bayes' rule. This justifies a time-independent denoiser for this setting. It does not justify dropping time for arbitrary replacement channels, state-dependent schedules, or corruption processes.

Check your understanding: For a general schedule, what weight remains in the forced-target estimator?

$-a'(t)$, since conditioning on a masked target cancels $1-a(t)$. If time is sampled from a density $g(t)$, the importance weight is $-a'(t)/g(t)$. The unweighted code above uses both the linear schedule and uniform time.

Where the derivation stops

The objective trains a denoiser. It does not make every parallel decoding rule exact. Continue to sampling and remasking to see where a correct one-position posterior becomes an approximate full-sequence transition.