U4.2 14 分鐘閱讀 2026年9月

U4.2 Masked Diffusion:ELBO 塌成加權的 Cross-Entropy

本篇重用M4.1進去就出不來:Absorbing State 與吸收時間·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy

上一篇那個一般的 ELBO,
在 absorbing 鏈上會塌成什麼?

把上一篇的 ELBO 算到底

上一篇留下一個一般的 bound:中間每一項 Lt1L_{t-1} 是 posterior q(xt1xt,x0)q(x_{t-1}\mid x_t,x_0) 與 model pθ(xt1xt)p_\theta(x_{t-1}\mid x_t) 之間的 KL。對一般的鏈,這個 KL 得逐項數值計算。但對 absorbing 鏈,它可以用手算完,而且結果簡單到會讓人懷疑是不是算錯了 [1, 2]。

先固定一個位置 \ell,記 αˉt\bar\alpha_t 為它到第 tt 步還沒被遮的機率(上一篇的符號)。位置 \ellxtx_t 裡只有兩種可能,分開看。

情形一:xtx_t^\ell 沒被遮。 上一篇說過,這時 posterior 是「xt1x_{t-1}^\ell 機率 1 是同一個原字」。而 x0x_0-parametrization 下的 model——只要我們讓網路遵守一個顯然的規則:看到沒被遮的位置,就照抄它——也給出同一個確定分佈。兩個相同的分佈,KL 是零。這個位置對 loss 沒有貢獻。

情形二:xt=x_t^\ell= [MASK] posterior 現在是一枚硬幣:

q(xt1xt,x0)={[MASK]機率 1αˉt11αˉt,x0機率 αˉt1αˉt1αˉt.q(x_{t-1}^\ell\mid x_t,x_0)=\begin{cases} \texttt{[MASK]} & \text{機率 } \dfrac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t},\\[6pt] x_0^\ell & \text{機率 } \dfrac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t}. \end{cases}

(分子分母都是「被遮的機率」:在 tt 被遮的前提下,t1t-1 時就已經被遮的機率是兩者之比。)model 這邊,pθ(xt1xt)=x~0q(xt1xt,x~0)pθ(x~0xt)p_\theta(x_{t-1}^\ell\mid x_t)=\sum_{\tilde x_0}q(x_{t-1}^\ell\mid x_t,\tilde x_0)\,p_\theta(\tilde x_0^\ell\mid x_t),而 q(xt,x~0)q(\cdot\mid x_t,\tilde x_0)任何 x~0\tilde x_0[MASK] 的機率都是同一個 1αˉt11αˉt\frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}。所以 model 也是同一枚硬幣:以同樣機率留在 [MASK],否則從 pθ(x~0xt)p_\theta(\tilde x_0^\ell\mid x_t) 抽一個字。

兩枚硬幣的「留在 [MASK]」那一面完全相同,KL 裡那一項相消;剩下「翻開」那一面:posterior 是集中在 x0x_0^\ell 的一個點,model 是 pθ(xt)p_\theta(\cdot\mid x_t)。點分佈對任意分佈的 KL 就是 logpθ(x0xt)-\log p_\theta(x_0^\ell\mid x_t)。乘上翻開的機率:

Lt1=αˉt1αˉt1αˉt  [logpθ(x0xt)](當 xt=[MASK]).L_{t-1}^{\ell}=\frac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t}\;\big[-\log p_\theta(x_0^\ell\mid x_t)\big]\qquad(\text{當 }x_t^\ell=\texttt{[MASK]}).

把所有位置、所有 tt 加起來,並注意 L0=logpθ(x0x1)L_0=-\log p_\theta(x_0\mid x_1) 剛好也是這個形式(代 t=1t=1αˉ0=1\bar\alpha_0=1 得權重 1),整個 ELBO 變成一條式子:

  LMDM=t=1T Extq(x0)[ αˉt1αˉt1αˉt: xt=[MASK]logpθ(x0xt)].  \boxed{\;\mathcal L_{\text{MDM}}=\sum_{t=1}^{T}\ \mathbb E_{x_t\sim q(\cdot\mid x_0)}\Bigg[\ \frac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t}\sum_{\ell:\ x_t^\ell=\texttt{[MASK]}}-\log p_\theta\big(x_0^\ell\mid x_t\big)\Bigg].\;}

讀一遍:抽一個 tt、按 αˉt\bar\alpha_t 隨機遮掉一些位置、對被遮的位置做 cross-entropy、乘上一個只跟 tt 有關的權重。 沒有 KL、沒有 posterior、沒有輔助 loss;上一篇 D3PM 額外加的 λ\lambda 項在這裡變成多餘的——ELBO 本身就是這個形狀 [1, 2]。

讓步數趨於無窮,這個和會變成一個積分。到了連續時間,「這一步」不再有意義,累積量是唯一有意義的東西,所以那條橫線會被省掉:從下面的展開框到 U5αt\alpha_t 指的都是這裡的 αˉt\bar\alpha_t

展開細節連續時間的極限,與線性 schedule 下的 1/t

TT\to\infty、把 αˉ\bar\alpha 看成 t[0,1]t\in[0,1] 上的一條光滑曲線 αt\alpha_t(從 α0=1\alpha_0=1 降到 α1=0\alpha_1=0),權重 αtΔαt1αtα˙t1αtdt\frac{\alpha_{t-\Delta}-\alpha_t}{1-\alpha_t}\to\frac{-\dot\alpha_t}{1-\alpha_t}\,dt,於是

LMDM=01α˙t1αt E[: xt=[MASK]logpθ(x0xt)]dt.\mathcal L_{\text{MDM}}=\int_0^1\frac{-\dot\alpha_t}{1-\alpha_t}\ \mathbb E\Bigg[\sum_{\ell:\ x_t^\ell=\texttt{[MASK]}}-\log p_\theta(x_0^\ell\mid x_t)\Bigg]dt .

α˙t<0\dot\alpha_t<0,所以權重是正的。)取線性 schedule αt=1t\alpha_t=1-t,權重是 1/t1/t;而 tt 時刻平均有 tLtL 個位置被遮,所以 1/t1/t 大致是在把「每個被遮位置的平均 cross-entropy」拉成等權——這是 MDLM [1] 與 MD4 [2] 都指出的事:這個 loss 本質上是一個對「遮罩比例」取平均的 cross-entropy。兩篇也都證明,在連續時間下換任何一條 αt\alpha_t 都只是把時間重新參數化,loss 的最小值不變;schedule 的選擇只影響訓練時 tt 的取樣密度,也就是 U2.4第四個旋鈕:訓練端的加權U3.5 把它和 SNR 的換算做完了)。

和 BERT 差在哪

算到這裡會覺得眼熟:對被遮的位置做 cross-entropy,就是 BERT [5] 的 masked language modeling。masked diffusion 和 BERT 的差別只有兩件事,而這兩件事把一個「表徵學習的預訓練目標」變成一個「合法的生成模型」:

  1. 遮罩比例是隨機的,而且要對所有比例都學會。 BERT 固定遮 15%;masked diffusion 的 tt 從 1 到 TT 都抽,對應遮 0% 到 100%。這是必要的,因為取樣時模型要從全部被遮一路走到全部翻開,中間每個比例都會遇到。MaskGIT [6] 在影像 token 上最早把這件事做成生成器。
  2. 每個 tt 有一個權重。 αˉt1αˉt1αˉt\frac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t} 是 ELBO 給的,它讓加總起來的 loss 是 logpθ(x0)-\log p_\theta(x_0) 的一個上界,所以訓出來的模型有 likelihood 可以報、可以和 autoregressive 模型比 perplexity。BERT 沒有這一層。

這個權重扮演的角色和 U1.2U3.5 談過的「對 noise level 加權」完全一樣——只是在這裡「noise level」變成了「遮掉的比例」。Kingma & Gao [7] 說連續世界的各種 diffusion loss 都是「ELBO 加上一個對 noise level 的權重函數」;masked diffusion 是這句話在離散世界最乾淨的例子。

展開細節一個有趣的副產品:網路其實不需要看 t

在 absorbing 鏈上,xtx_t 本身就洩漏了 tt 的資訊——被遮的比例大致是 1αˉt1-\bar\alpha_t。更強的是:p(x0xt)p(x_0^\ell\mid x_t) 這個條件分佈根本不依賴 tt,它只依賴「哪些位置被遮、沒被遮的位置是什麼字」;因為給定沒被遮的字,x0x_0^\ell 的條件分佈就是資料分佈的一個條件分佈,跟你是怎麼走到這個遮罩狀態的無關。Ou et al. [3] 把這件事說清楚:absorbing diffusion 學的就是乾淨資料的所有條件分佈 pdata(xx未遮)p_{\text{data}}(x^\ell\mid x^{\text{未遮}}),網路可以拿掉時間輸入。這也解釋了它為什麼和 BERT 這麼像——BERT 學的正是這些條件分佈的一部分。

補充為什麼這個權重長成 1/t,以及它在實務上常被改掉

線性 schedule 下 αˉt=1t\bar\alpha_t=1-t,權重 α˙t1αt=1t\frac{-\dot\alpha_t}{1-\alpha_t}=\frac1t——tt 小(遮得少)的時候權重很大。直覺是:遮得少的時候每一個被遮的位置都很好猜,但 ELBO 要求那些容易的位置也要猜得很準。

實務上很多實作會把這個權重改掉(例如換成常數、或對 tt 用別的分布),代價是失去 likelihood 的界——訓出來的東西還是可以生成,但不能再報 perplexity。這個取捨和 U1.2 說「DDPM 故意丟掉理論加權」是同一件事。

取樣:從全 [MASK] 出發

有了 pθ(x0xt)p_\theta(x_0\mid x_t),取樣就是把上一篇的 posterior 反著走。從 xT=x_T=[MASK] 出發,每一步 tt1t\to t-1,對每個仍是 [MASK] 的位置擲上面那枚硬幣:

  • 以機率 αˉt1αˉt1αˉt\frac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t} 翻開:從 pθ(xt)p_\theta(\cdot\mid x_t) 抽一個字填進去;
  • 否則留在 [MASK],下一步再說。

已經翻開的位置不再改變——這不是額外的規定,是 absorbing 鏈的 posterior 說的(沒被遮的 xtx_t^\ellxt1x_{t-1}^\ell 機率 1 是同一個字)。上一篇 Q1 提過的跳步在這裡直接可用:把 tt1t\to t-1 換成 ttkt\to t-k,硬幣的機率換成 αˉtkαˉt1αˉt\frac{\bar\alpha_{t-k}-\bar\alpha_t}{1-\bar\alpha_t},其餘不變。所以 masked diffusion 可以訓練時用 T=1000T=1000、取樣時只走 8 步或 16 步。

互動 demo:從全 [MASK] 出發,每步翻開一部分。 資料是一個 toy 文法(長度 16、字典 A B ( )、A 與 B 交替且括號配對)。按「下一步」看每一步翻開哪些格子——機率只跟 tt 有關,和內容無關。把步數從 16 降到 1,每步同時翻開的格子越來越多,最後一步就把 16 格一起填完。點一個還被遮的格子看它的分布:翻開時是從那幾根柱子抽的,不是取 argmax。最後按 BERT 對照:BERT 只在一個固定的遮罩比例上訓練、權重是常數;masked diffusion 要對所有比例都學會,而且各比例的權重不一樣。

課堂提問Q1

取樣時每一步翻開的字是從 pθ(xt)p_\theta(\cdot\mid x_t) 出來的。如果改成每次都取 argmax(最可能的字),會怎樣?

先想一想,再展開看整理後的答案

會失去多樣性,而且會失去正確性。從同一個全 [MASK] 出發、每步都 argmax,整個程序變成確定性的(除了硬幣決定翻開哪些位置),所有樣本會擠向同一批「最安全」的句子。更根本的是,pθ(x0xt)p_\theta(x_0^\ell\mid x_t) 是一個 conditional distribution,argmax 只在這個分佈非常尖的時候才接近「從中抽樣」;遮罩比例大時它通常很平,argmax 偏得很遠。

不過這個問題有一個實務上的變種值得留著:MaskGIT [6] 的做法是每步先對所有被遮位置抽樣,再只保留信心最高的那幾個,其餘遮回去。那不是 argmax,而是在「哪些位置翻開」上做選擇——它其實在偷偷處理下一篇要談的問題。

先消化一下

想一想

在 absorbing 鏈的 ELBO 裡,沒被遮的位置對 loss 的貢獻是零,原因是:

想一想

Masked diffusion 的 loss 與 BERT 的 masked language modeling 相比,多出來的兩件事是:

想一想

取樣時「已經翻開的 token 之後不再改變」這件事的來源是:

想一想

線性 schedule αt=1t\alpha_t=1-t 下連續時間的權重是 1/t1/t。它在做的事最接近:

參考文獻

  1. Sahoo, S. S., Arriola, M., Schiff, Y., Gokaslan, A., Marroquin, E., Chiu, J. T., Rush, A., Kuleshov, V. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.(MDLM:ELBO 化簡為加權 cross-entropy、SUBS parametrization、連續時間極限與 schedule 不變性。)
  2. Shi, J., Han, K., Wang, Z., Doucet, A., Titsias, M. K. Simplified and Generalized Masked Diffusion for Discrete Data. NeurIPS 2024.(MD4:同一個化簡的獨立推導與推廣。)
  3. Ou, J., Nie, S., Xue, K., Zhu, F., Sun, J., Li, Z., Li, C. Your Absorbing Discrete Diffusion Secretly Models the Conditional Distributions of Clean Data. ICLR 2025.(網路不需要時間輸入。)
  4. Austin, J., Johnson, D. D., Ho, J., Tarlow, D., van den Berg, R. Structured Denoising Diffusion Models in Discrete State-Spaces. NeurIPS 2021.
  5. Devlin, J., Chang, M.-W., Lee, K., Toutanova, K. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. NAACL 2019.
  6. Chang, H., Zhang, H., Jiang, L., Liu, C., Freeman, W. T. MaskGIT: Masked Generative Image Transformer. CVPR 2022.(隨機遮罩比例的生成器;信心排序的 unmask 策略。)
  7. Kingma, D. P., Gao, R. Understanding Diffusion Objectives as the ELBO with Simple Data Augmentation. NeurIPS 2023.(各種 diffusion loss 都是 ELBO 加 noise level 權重。)