學完這篇,你能
對 absorbing 鏈把 ELBO 的每一項算出來,說出為什麼沒被遮的位置貢獻為零、被遮的位置只剩一個 cross-entropy。
寫出 masked diffusion 的 loss,指出它與 BERT 的 masked language modeling 差在哪兩件事。
描述取樣程序(從全 [MASK] 出發、每步 unmask 一部分),並說出為什麼已 unmask 的 token 之後不再改變。
本篇重用 M4.1 進去就出不來:Absorbing State 與吸收時間· M5.1 用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy
上一篇那個一般的 ELBO,
在 absorbing 鏈上會塌成什麼?
把上一篇的 ELBO 算到底
上一篇留下一個一般的 bound:中間每一項 L t − 1 L_{t-1} L t − 1 是 posterior q ( x t − 1 ∣ x t , x 0 ) q(x_{t-1}\mid x_t,x_0) q ( x t − 1 ∣ x t , x 0 ) 與 model p θ ( x t − 1 ∣ x t ) p_\theta(x_{t-1}\mid x_t) p θ ( x t − 1 ∣ x t ) 之間的 KL。對一般的鏈,這個 KL 得逐項數值計算。但對 absorbing 鏈,它可以用手算完 ,而且結果簡單到會讓人懷疑是不是算錯了 [1, 2]。
先固定一個位置 ℓ \ell ℓ ,記 α ˉ t \bar\alpha_t α ˉ t 為它到第 t t t 步還沒被遮的機率(上一篇的符號)。位置 ℓ \ell ℓ 在 x t x_t x t 裡只有兩種可能,分開看。
情形一:x t ℓ x_t^\ell x t ℓ 沒被遮。 上一篇說過,這時 posterior 是「x t − 1 ℓ x_{t-1}^\ell x t − 1 ℓ 機率 1 是同一個原字」。而 x 0 x_0 x 0 -parametrization 下的 model——只要我們讓網路遵守一個顯然的規則:看到沒被遮的位置,就照抄它 ——也給出同一個確定分佈。兩個相同的分佈,KL 是零。這個位置對 loss 沒有貢獻。
情形二:x t ℓ = x_t^\ell= x t ℓ = [MASK]。 posterior 現在是一枚硬幣:
q ( x t − 1 ℓ ∣ x t , x 0 ) = { [MASK] 機率 1 − α ˉ t − 1 1 − α ˉ t , x 0 ℓ 機率 α ˉ t − 1 − α ˉ t 1 − α ˉ 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} q ( x t − 1 ℓ ∣ x t , x 0 ) = ⎩ ⎨ ⎧ [MASK] x 0 ℓ 機率 1 − α ˉ t 1 − α ˉ t − 1 , 機率 1 − α ˉ t α ˉ t − 1 − α ˉ t .
(分子分母都是「被遮的機率」:在 t t t 被遮的前提下,t − 1 t-1 t − 1 時就已經被遮的機率是兩者之比。)model 這邊,p θ ( x t − 1 ℓ ∣ x t ) = ∑ x ~ 0 q ( x t − 1 ℓ ∣ x t , x ~ 0 ) p θ ( x ~ 0 ℓ ∣ x t ) 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) p θ ( x t − 1 ℓ ∣ x t ) = ∑ x ~ 0 q ( x t − 1 ℓ ∣ x t , x ~ 0 ) p θ ( x ~ 0 ℓ ∣ x t ) ,而 q ( ⋅ ∣ x t , x ~ 0 ) q(\cdot\mid x_t,\tilde x_0) q ( ⋅ ∣ x t , x ~ 0 ) 對任何 x ~ 0 \tilde x_0 x ~ 0 給 [MASK] 的機率都是同一個 1 − α ˉ t − 1 1 − α ˉ t \frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t} 1 − α ˉ t 1 − α ˉ t − 1 。所以 model 也是同一枚硬幣:以同樣機率留在 [MASK],否則從 p θ ( x ~ 0 ℓ ∣ x t ) p_\theta(\tilde x_0^\ell\mid x_t) p θ ( x ~ 0 ℓ ∣ x t ) 抽一個字。
兩枚硬幣的「留在 [MASK]」那一面完全相同,KL 裡那一項相消;剩下「翻開」那一面:posterior 是集中在 x 0 ℓ x_0^\ell x 0 ℓ 的一個點,model 是 p θ ( ⋅ ∣ x t ) p_\theta(\cdot\mid x_t) p θ ( ⋅ ∣ x t ) 。點分佈對任意分佈的 KL 就是 − log p θ ( x 0 ℓ ∣ x t ) -\log p_\theta(x_0^\ell\mid x_t) − log p θ ( x 0 ℓ ∣ x t ) 。乘上翻開的機率:
L t − 1 ℓ = α ˉ t − 1 − α ˉ t 1 − α ˉ t [ − log p θ ( x 0 ℓ ∣ x t ) ] ( 當 x t ℓ = [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]}). L t − 1 ℓ = 1 − α ˉ t α ˉ t − 1 − α ˉ t [ − log p θ ( x 0 ℓ ∣ x t ) ] ( 當 x t ℓ = [MASK] ) .
把所有位置、所有 t t t 加起來,並注意 L 0 = − log p θ ( x 0 ∣ x 1 ) L_0=-\log p_\theta(x_0\mid x_1) L 0 = − log p θ ( x 0 ∣ x 1 ) 剛好也是這個形式(代 t = 1 t=1 t = 1 、α ˉ 0 = 1 \bar\alpha_0=1 α ˉ 0 = 1 得權重 1),整個 ELBO 變成一條式子:
L MDM = ∑ t = 1 T E x t ∼ q ( ⋅ ∣ x 0 ) [ α ˉ t − 1 − α ˉ t 1 − α ˉ t ∑ ℓ : x t ℓ = [MASK] − log p θ ( x 0 ℓ ∣ x t ) ] . \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].\;} L MDM = t = 1 ∑ T E x t ∼ q ( ⋅ ∣ x 0 ) [ 1 − α ˉ t α ˉ t − 1 − α ˉ t ℓ : x t ℓ = [MASK] ∑ − log p θ ( x 0 ℓ ∣ x t ) ] .
讀一遍:抽一個 t t t 、按 α ˉ t \bar\alpha_t α ˉ t 隨機遮掉一些位置、對被遮的位置做 cross-entropy、乘上一個只跟 t t t 有關的權重。 沒有 KL、沒有 posterior、沒有輔助 loss;上一篇 D3PM 額外加的 λ \lambda λ 項在這裡變成多餘的——ELBO 本身就是這個形狀 [1, 2]。
示意圖 · 規劃中 ELBO 為什麼會塌
位置 :本節之後。檔名 :week-6/imgs/w6-2-1.svg(16:9,三欄)。
左欄 :D3PM 的 ELBO 三種項 L 0 L_0 L 0 、∑ L t − 1 \sum L_{t-1} ∑ L t − 1 、L T L_T L T 疊成三個方塊,L T L_T L T 標「≈ 0 \approx0 ≈ 0 ,無參數」淡出。
中欄 :一個 absorbing 位置的兩枚硬幣並排——左硬幣 posterior、右硬幣 model。兩枚硬幣「留在 [MASK]」的那一面畫成同色、中間一個等號並劃掉(相消);「翻開」那一面:posterior 是一根集中在 x 0 ℓ x_0^\ell x 0 ℓ 的柱子,model 是 K K K 根柱子的分佈,兩者之間標 − log p θ ( x 0 ℓ ∣ x t ) -\log p_\theta(x_0^\ell\mid x_t) − log p θ ( x 0 ℓ ∣ x t ) 。上方一行小字:「沒被遮的位置:兩邊都是確定的 → KL = 0 =0 = 0 」。
右欄 :只剩一條式子 ∑ t w t ∑ ℓ ∈ masked − log p θ ( x 0 ℓ ∣ x t ) \sum_t w_t\sum_{\ell\in\text{masked}}-\log p_\theta(x_0^\ell\mid x_t) ∑ t w t ∑ ℓ ∈ masked − log p θ ( x 0 ℓ ∣ x t ) ,w t w_t w t 用醒目色框起,旁邊小字「只跟 t t t 有關」。
圖說(.mdx 內) :「圖:ELBO 為什麼會塌。 沒被遮的位置兩邊都是確定的,KL 為零;被遮的位置兩邊是同一枚硬幣,『留在 [MASK]』那一面相消,只剩翻開時的 cross-entropy。」
術語一律英文(posterior、model、cross-entropy、masked);配色沿用 U1.0 圖 1 的基準色。
讓步數趨於無窮,這個和會變成一個積分。到了連續時間,「這一步」不再有意義,累積量是唯一有意義的東西,所以那條橫線會被省掉 :從下面的展開框到 U5 ,α t \alpha_t α t 指的都是這裡的 α ˉ t \bar\alpha_t α ˉ t 。
展開細節 連續時間的極限,與線性 schedule 下的 1/t 讓 T → ∞ T\to\infty T → ∞ 、把 α ˉ \bar\alpha α ˉ 看成 t ∈ [ 0 , 1 ] t\in[0,1] t ∈ [ 0 , 1 ] 上的一條光滑曲線 α t \alpha_t α t (從 α 0 = 1 \alpha_0=1 α 0 = 1 降到 α 1 = 0 \alpha_1=0 α 1 = 0 ),權重 α t − Δ − α t 1 − α t → − α ˙ t 1 − α t d t \frac{\alpha_{t-\Delta}-\alpha_t}{1-\alpha_t}\to\frac{-\dot\alpha_t}{1-\alpha_t}\,dt 1 − α t α t − Δ − α t → 1 − α t − α ˙ t d t ,於是
L MDM = ∫ 0 1 − α ˙ t 1 − α t E [ ∑ ℓ : x t ℓ = [MASK] − log p θ ( x 0 ℓ ∣ x t ) ] d t . \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 . L MDM = ∫ 0 1 1 − α t − α ˙ t E [ ℓ : x t ℓ = [MASK] ∑ − log p θ ( x 0 ℓ ∣ x t ) ] d t . (α ˙ t < 0 \dot\alpha_t<0 α ˙ t < 0 ,所以權重是正的。)取線性 schedule α t = 1 − t \alpha_t=1-t α t = 1 − t ,權重是 1 / t 1/t 1/ t ;而 t t t 時刻平均有 t L tL t L 個位置被遮,所以 1 / t 1/t 1/ t 大致是在把「每個被遮位置的平均 cross-entropy」拉成等權——這是 MDLM [1] 與 MD4 [2] 都指出的事:這個 loss 本質上是一個對「遮罩比例」取平均的 cross-entropy。兩篇也都證明,在連續時間下換任何一條 α t \alpha_t α t 都只是把時間重新參數化,loss 的最小值不變;schedule 的選擇只影響訓練時 t t t 的取樣密度,也就是 U2.4 那第四個旋鈕:訓練端的加權 (U3.5 把它和 SNR 的換算做完了)。
和 BERT 差在哪
算到這裡會覺得眼熟:對被遮的位置做 cross-entropy ,就是 BERT [5] 的 masked language modeling。masked diffusion 和 BERT 的差別只有兩件事,而這兩件事把一個「表徵學習的預訓練目標」變成一個「合法的生成模型」:
遮罩比例是隨機的,而且要對所有比例都學會。 BERT 固定遮 15%;masked diffusion 的 t t t 從 1 到 T T T 都抽,對應遮 0% 到 100%。這是必要的,因為取樣時模型要從全部被遮一路走到全部翻開,中間每個比例都會遇到。MaskGIT [6] 在影像 token 上最早把這件事做成生成器。
每個 t t t 有一個權重。 α ˉ t − 1 − α ˉ t 1 − α ˉ t \frac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t} 1 − α ˉ t α ˉ t − 1 − α ˉ t 是 ELBO 給的,它讓加總起來的 loss 是 − log p θ ( x 0 ) -\log p_\theta(x_0) − log p θ ( x 0 ) 的一個上界,所以訓出來的模型有 likelihood 可以報、可以和 autoregressive 模型比 perplexity。BERT 沒有這一層。
這個權重扮演的角色和 U1.2 、U3.5 談過的「對 noise level 加權」完全一樣——只是在這裡「noise level」變成了「遮掉的比例」。Kingma & Gao [7] 說連續世界的各種 diffusion loss 都是「ELBO 加上一個對 noise level 的權重函數」;masked diffusion 是這句話在離散世界最乾淨的例子。
展開細節 一個有趣的副產品:網路其實不需要看 t 在 absorbing 鏈上,x t x_t x t 本身就洩漏了 t t t 的資訊——被遮的比例大致是 1 − α ˉ t 1-\bar\alpha_t 1 − α ˉ t 。更強的是:p ( x 0 ℓ ∣ x t ) p(x_0^\ell\mid x_t) p ( x 0 ℓ ∣ x t ) 這個條件分佈根本不依賴 t t t ,它只依賴「哪些位置被遮、沒被遮的位置是什麼字」;因為給定沒被遮的字,x 0 ℓ x_0^\ell x 0 ℓ 的條件分佈就是資料分佈的一個條件分佈,跟你是怎麼走到這個遮罩狀態的無關。Ou et al. [3] 把這件事說清楚:absorbing diffusion 學的就是乾淨資料的所有條件分佈 p data ( x ℓ ∣ x 未遮 ) p_{\text{data}}(x^\ell\mid x^{\text{未遮}}) p data ( x ℓ ∣ x 未遮 ) ,網路可以拿掉時間輸入。這也解釋了它為什麼和 BERT 這麼像——BERT 學的正是這些條件分佈的一部分。
補充 為什麼這個權重長成 1/t,以及它在實務上常被改掉 線性 schedule 下 α ˉ t = 1 − t \bar\alpha_t=1-t α ˉ t = 1 − t ,權重 − α ˙ t 1 − α t = 1 t \frac{-\dot\alpha_t}{1-\alpha_t}=\frac1t 1 − α t − α ˙ t = t 1 ——t t t 小(遮得少)的時候權重很大。直覺是:遮得少的時候每一個被遮的位置都很好猜,但 ELBO 要求那些容易的位置也要猜得很準。
實務上很多實作會把這個權重改掉(例如換成常數、或對 t t t 用別的分布),代價是失去 likelihood 的界 ——訓出來的東西還是可以生成,但不能再報 perplexity。這個取捨和 U1.2 說「DDPM 故意丟掉理論加權」是同一件事。
取樣:從全 [MASK] 出發
有了 p θ ( x 0 ∣ x t ) p_\theta(x_0\mid x_t) p θ ( x 0 ∣ x t ) ,取樣就是把上一篇的 posterior 反著走。從 x T = x_T= x T = 全 [MASK] 出發,每一步 t → t − 1 t\to t-1 t → t − 1 ,對每個仍是 [MASK] 的位置擲上面那枚硬幣:
以機率 α ˉ t − 1 − α ˉ t 1 − α ˉ t \frac{\bar\alpha_{t-1}-\bar\alpha_t}{1-\bar\alpha_t} 1 − α ˉ t α ˉ t − 1 − α ˉ t 翻開 :從 p θ ( ⋅ ∣ x t ) p_\theta(\cdot\mid x_t) p θ ( ⋅ ∣ x t ) 抽一個字填進去;
否則留在 [MASK],下一步再說。
已經翻開的位置不再改變 ——這不是額外的規定,是 absorbing 鏈的 posterior 說的(沒被遮的 x t ℓ x_t^\ell x t ℓ ,x t − 1 ℓ x_{t-1}^\ell x t − 1 ℓ 機率 1 是同一個字)。上一篇 Q1 提過的跳步在這裡直接可用:把 t → t − 1 t\to t-1 t → t − 1 換成 t → t − k t\to t-k t → t − k ,硬幣的機率換成 α ˉ t − k − α ˉ t 1 − α ˉ t \frac{\bar\alpha_{t-k}-\bar\alpha_t}{1-\bar\alpha_t} 1 − α ˉ t α ˉ t − k − α ˉ t ,其餘不變。所以 masked diffusion 可以訓練時用 T = 1000 T=1000 T = 1000 、取樣時只走 8 步或 16 步。
互動 demo:從全 [MASK] 出發,每步翻開一部分。 資料是一個 toy 文法(長度 16、字典 A B ( )、A 與 B 交替且括號配對)。按「下一步」看每一步翻開哪些格子——機率只跟 t t t 有關,和內容無關 。把步數從 16 降到 1,每步同時翻開的格子越來越多,最後一步就把 16 格一起填完。點一個還被遮的格子看它的分布:翻開時是從那幾根柱子抽的 ,不是取 argmax。最後按 BERT 對照:BERT 只在一個固定的遮罩比例上訓練、權重是常數;masked diffusion 要對所有比例都學會,而且各比例的權重不一樣。
課堂提問 Q1
取樣時每一步翻開的字是從 p θ ( ⋅ ∣ x t ) p_\theta(\cdot\mid x_t) p θ ( ⋅ ∣ x t ) 抽 出來的。如果改成每次都取 argmax(最可能的字),會怎樣?
先想一想,再展開看整理後的答案 會失去多樣性,而且會失去正確性。從同一個全 [MASK] 出發、每步都 argmax,整個程序變成確定性的(除了硬幣決定翻開哪些位置),所有樣本會擠向同一批「最安全」的句子。更根本的是,p θ ( x 0 ℓ ∣ x t ) p_\theta(x_0^\ell\mid x_t) p θ ( x 0 ℓ ∣ x t ) 是一個 conditional distribution,argmax 只在這個分佈非常尖的時候才接近「從中抽樣」;遮罩比例大時它通常很平,argmax 偏得很遠。
不過這個問題有一個實務上的變種值得留著:MaskGIT [6] 的做法是每步先對所有被遮位置抽樣,再只保留信心最高的那幾個 ,其餘遮回去。那不是 argmax,而是在「哪些位置翻開」上做選擇——它其實在偷偷處理下一篇要談的問題。
先消化一下
參考文獻
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 不變性。)
Shi, J., Han, K., Wang, Z., Doucet, A., Titsias, M. K. Simplified and Generalized Masked Diffusion for Discrete Data. NeurIPS 2024.(MD4:同一個化簡的獨立推導與推廣。)
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.(網路不需要時間輸入。)
Austin, J., Johnson, D. D., Ho, J., Tarlow, D., van den Berg, R. Structured Denoising Diffusion Models in Discrete State-Spaces. NeurIPS 2021.
Devlin, J., Chang, M.-W., Lee, K., Toutanova, K. BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding. NAACL 2019.
Chang, H., Zhang, H., Jiang, L., Liu, C., Freeman, W. T. MaskGIT: Masked Generative Image Transformer. CVPR 2022.(隨機遮罩比例的生成器;信心排序的 unmask 策略。)
Kingma, D. P., Gao, R. Understanding Diffusion Objectives as the ELBO with Simple Data Augmentation. NeurIPS 2023.(各種 diffusion loss 都是 ELBO 加 noise level 權重。)