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

U4.1 D3PM:Transition Matrix、Closed Form 與 x₀-Parametrization

本篇重用M4.0明天的天氣只看今天:Markov Chain 與 Transition Matrix·M4.1進去就出不來:Absorbing State 與吸收時間·M1.3快篩陽性:Bayes 與 Posterior·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy

上一篇挑好了兩條鏈。
那「q(xtx0)q(x_t\mid x_0) 有 closed form」在離散世界是什麼意思?

先只看一個 token

一句話有很多 token,但 forward process 對每個位置是獨立地加噪聲,所以先只盯著一個位置。這個位置的狀態 x{1,,K}x\in\{1,\dots,K\},寫成 one-hot 的列向量(row vector)x{0,1}Kx\in\{0,1\}^{K}

一條 Markov chain 每一步的規則是一個 K×KK\times Ktransition matrix QtQ_t

[Qt]ij=q(xt=jxt1=i),[Q_t]_{ij}=q(x_t=j\mid x_{t-1}=i),

ii 列是「從狀態 ii 出發,下一步落在各狀態的機率」,每列和為 1。用 one-hot 列向量寫,q(xtxt1)=Cat(xt1Qt)q(x_t\mid x_{t-1})=\mathrm{Cat}(x_{t-1}Q_t)——xt1x_{t-1} 挑出 QtQ_t 的那一列,那一列就是 xtx_t 的分佈。

tt 步就是連乘:

q(xtx0)=Cat(x0Qˉt),Qˉt=Q1Q2Qt.q(x_t\mid x_0)=\mathrm{Cat}\big(x_0\,\bar Q_t\big),\qquad \bar Q_t=Q_1Q_2\cdots Q_t .

這一行是整個單元的地基。U1.1 說「forward 是我們自己定的,所以任何 ttxtx_t 可以一步抽出來」——在離散世界裡,這句話就是「Qˉt\bar Q_t 有 closed form」。一般的矩陣連乘沒有 closed form,但我們選的兩條鏈剛好都有。

這兩條鏈,Qˉt\bar Q_t 真的算得出來

Uniform。 每步以機率 βt\beta_t 把 token 換成均勻隨機的一個:

Qt=(1βt)I+βt1K11 ⁣.Q_t=(1-\beta_t)\,I+\beta_t\,\tfrac1K\mathbb 1\mathbb 1^{\!\top}.

J=1K11 ⁣J=\tfrac1K\mathbb 1\mathbb 1^{\!\top}(每個元素都是 1/K1/K 的矩陣),它滿足 J2=JJ^2=J。形如 aI+bJaI+bJa+b=1a+b=1 的矩陣相乘還是這個形狀:(aI+bJ)(cI+dJ)=acI+(1ac)J(aI+bJ)(cI+dJ)=ac\,I+(1-ac)\,J。於是

Qˉt=αˉtI+(1αˉt)J,αˉt=st(1βs).\bar Q_t=\bar\alpha_t\,I+(1-\bar\alpha_t)\,J,\qquad \bar\alpha_t=\prod_{s\le t}(1-\beta_s).

讀法:到了第 tt 步,token 以機率 αˉt\bar\alpha_t 從沒被動過、還是原字;以機率 1αˉt1-\bar\alpha_t 已經被換成一個均勻隨機的字。αˉt0\bar\alpha_t\to0QˉtJ\bar Q_t\to J,每個位置獨立均勻——這就是 uniform 鏈的終點。

Absorbing。 字典加一個第 K+1K{+}1 個狀態 [MASK],記它的 one-hot 為 eme_m。每步以機率 βt\beta_t 把 token 變成 [MASK][MASK] 自己不動:

Qt=(1βt)I+βt1em ⁣.Q_t=(1-\beta_t)\,I+\beta_t\,\mathbb 1 e_m^{\!\top}.

(檢查第 mm 列:(1βt)em ⁣+βtem ⁣=em ⁣(1-\beta_t)e_m^{\!\top}+\beta_t e_m^{\!\top}=e_m^{\!\top}[MASK] 一定留在 [MASK]。)同樣的代數,因為 em ⁣1=1e_m^{\!\top}\mathbb 1=1

Qˉt=αˉtI+(1αˉt)1em ⁣.\bar Q_t=\bar\alpha_t\,I+(1-\bar\alpha_t)\,\mathbb 1 e_m^{\!\top}.

讀法:到了第 tt 步,token 以機率 αˉt\bar\alpha_t 還是原字,以機率 1αˉt1-\bar\alpha_t 已經是 [MASK]沒有第三種可能。 這就是上一篇說的「沒被遮的位置一定是原字」,現在它是一個公式。

兩條鏈的 αˉt\bar\alpha_t 扮演的角色,和 U1.1αˉt\sqrt{\bar\alpha_t}x0x_0 的係數一樣:還剩多少原始訊號。這是我們刻意沿用同一個符號的原因。

符號這裡的 alpha bar 和連續世界的差一個平方根

連續世界寫 xt=αˉtx0+σtϵx_t=\sqrt{\bar\alpha_t}\,x_0+\sigma_t\epsilon,係數是 αˉt\sqrt{\bar\alpha_t};這裡寫 Qˉt=αˉtI+(1αˉt)()\bar Q_t=\bar\alpha_t I+(1-\bar\alpha_t)(\cdot),係數是 αˉt\bar\alpha_t 本身。

差別來自「係數」的意思不同:連續世界那個是振幅(能量是它的平方),這裡是機率(token 沒被動過的機率)。兩邊都用 αˉt\bar\alpha_t 是為了讓「還剩多少原始訊號」這個直覺可以直接搬過來,但代進式子的時候不要記錯。

互動 demo:Qˉt\bar Q_t 的熱圖隨 tt 演化。tt:uniform 的對角線一路淡下去、其他格一起亮起來,最後整張變成一片均勻;absorbing 的質量整批流進 [MASK] 那一欄,而最後一列從頭到尾都是 [MASK] 自己。切到「QtQ_t(只走一步)」會看到單步矩陣幾乎就是 identity——是連乘 tt 次才把質量搬走的。另外對照兩條鏈的對角線:uniform 的「還是原字」比 αˉt\bar\alpha_t 多一點(多了「被換掉但剛好換回自己」)。

Posterior:多知道 x0x_0,一切都好算

U1.2 有一個關鍵步驟:反向那一步 q(xt1xt)q(x_{t-1}\mid x_t) 不好算,但多給 x0x_0 之後 q(xt1xt,x0)q(x_{t-1}\mid x_t,x_0) 是一個有 closed form 的 Gaussian。同一件事在這裡更簡單,因為一切都是有限個數字。Bayes 定理:

q(xt1xt,x0)    q(xtxt1)  q(xt1x0)  =  (xtQt ⁣)(x0Qˉt1),q(x_{t-1}\mid x_t,x_0)\;\propto\;q(x_t\mid x_{t-1})\;q(x_{t-1}\mid x_0) \;=\;\big(x_t Q_t^{\!\top}\big)\odot\big(x_0\bar Q_{t-1}\big),

\odot 是逐元素相乘,兩個都是長度 KK 的向量(分別是「哪些 xt1x_{t-1} 走一步能到 xtx_t」和「從 x0x_0t1t-1 步會到哪些 xt1x_{t-1}」),乘起來再正規化,就是 xt1x_{t-1} 的類別分佈。

對 absorbing 鏈這個 posterior 特別直白:若 xtx_t 沒被遮,那 xt1x_{t-1} 一定也是同一個原字(機率 1);若 xt=x_t= [MASK],那 xt1x_{t-1} 以某個機率是 [MASK]、否則就是 x0x_0已知 x0x_0 之後,反向每一步都只是「要不要把這個 [MASK] 換回 x0x_0」的一個硬幣。

那 loss 長什麼樣?

現在寫訓練目標。U1.2 為了盡快到達那五行程式碼,直接用「預測 x0x_0 的 MSE」當 loss,把 DDPM 原本的 variational bound 放在一旁。在離散世界裡值得把這個 bound 攤開看一次,因為它的每一項都是有限個數字之間的 KL,沒有任何近似。

Reverse model 是 pθ(xt1xt)p_\theta(x_{t-1}\mid x_t),起點 p(xT)p(x_T) 是鏈的終點分佈。DDPM [4] 的 bound 一字不改地成立 [1, 2]:

logpθ(x0)    Eq[logpθ(x0x1)]L0+t=2TEqKL(q(xt1xt,x0)pθ(xt1xt))Lt1+KL(q(xTx0)p(xT))LT.-\log p_\theta(x_0)\;\le\; \underbrace{\mathbb E_q\big[-\log p_\theta(x_0\mid x_1)\big]}_{L_0} +\sum_{t=2}^{T}\underbrace{\mathbb E_q\,\mathrm{KL}\big(q(x_{t-1}\mid x_t,x_0)\,\big\|\,p_\theta(x_{t-1}\mid x_t)\big)}_{L_{t-1}} +\underbrace{\mathrm{KL}\big(q(x_T\mid x_0)\,\big\|\,p(x_T)\big)}_{L_T}.

LTL_T 沒有參數,而且對我們選的兩條鏈它趨近 0(αˉT0\bar\alpha_T\approx0q(xTx0)q(x_T\mid x_0) 就是終點分佈)。L0L_0 是一個普通的 cross-entropy。中間每一項 Lt1L_{t-1} 是兩個 KK 維類別分佈的 KL——左邊是那條 Bayes 算出來的 posterior q(xt1xt,x0)q(x_{t-1}\mid x_t,x_0),右邊是網路。

展開細節推導的骨架:為什麼 bound 長這樣(與 DDPM 逐字相同)

logpθ(x0)=logpθ(x0:T)dx1:T\log p_\theta(x_0)=\log\int p_\theta(x_{0:T})\,dx_{1:T} 出發,用 q(x1:Tx0)q(x_{1:T}\mid x_0) 當 proposal 套 Jensen 不等式:

logpθ(x0)Eq(x1:Tx0)[logq(x1:Tx0)pθ(x0:T)].-\log p_\theta(x_0)\le\mathbb E_{q(x_{1:T}\mid x_0)}\Big[\log\frac{q(x_{1:T}\mid x_0)}{p_\theta(x_{0:T})}\Big].

q(x1:Tx0)=tq(xtxt1)q(x_{1:T}\mid x_0)=\prod_t q(x_t\mid x_{t-1}) 用 Bayes 改寫成 q(xTx0)t2q(xt1xt,x0)q(x_T\mid x_0)\prod_{t\ge2}q(x_{t-1}\mid x_t,x_0)pθ(x0:T)=p(xT)tpθ(xt1xt)p_\theta(x_{0:T})=p(x_T)\prod_t p_\theta(x_{t-1}\mid x_t),逐項配對,就得到上面三種項。在連續世界裡積分換成離散世界的求和,其餘一步都不用改;這就是為什麼說「bound 一字不改地成立」。

網路該輸出什麼

還有一個設計決定:pθ(xt1xt)p_\theta(x_{t-1}\mid x_t) 要怎麼參數化?

最直接的想法是讓網路直接輸出 xt1x_{t-1}KK 個 logits。D3PM [1] 的選擇不同,而且是這一篇最值得記的一句話:讓網路預測 x0x_0,再用已知的 posterior 把它折回 xt1x_{t-1}

pθ(xt1xt)=x~0q(xt1xt,x~0)  pθ(x~0xt).p_\theta(x_{t-1}\mid x_t)=\sum_{\tilde x_0}q(x_{t-1}\mid x_t,\tilde x_0)\;p_\theta(\tilde x_0\mid x_t).

網路輸出的是 pθ(x~0xt)p_\theta(\tilde x_0\mid x_t) 的 logits(每個位置 KK 個數);q(xt1xt,x~0)q(x_{t-1}\mid x_t,\tilde x_0) 是上面那個 closed form,不用學。這叫 x0x_0-parametrization,它就是 U1.2x0x_0-prediction 的離散翻譯——在那裡網路預測 x^0\hat x_0、再用 q(xt1xt,x^0)q(x_{t-1}\mid x_t,\hat x_0) 的 Gaussian posterior 往回走一步;在這裡完全一樣,只是「預測 x0x_0」從輸出一個向量變成輸出一個類別分佈。

課堂提問Q1

U1.2x0x_0-、ϵ\epsilon-、vv-prediction 是三個線性可互換的座標,實務上常選 ϵ\epsilon。離散世界裡為什麼只剩 x0x_0-parametrization 這一種選擇?

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

因為那三個座標可以互換,靠的是 xt=αˉtx0+σtϵx_t=\sqrt{\bar\alpha_t}\,x_0+\sigma_t\epsilon 這條 affine 關係——知道 xtx_t 與任一個,就能解出另外兩個。離散世界沒有這條關係:xtx_t 是從 x0x_0 出來的一個類別,沒有一個叫「ϵ\epsilon」的連續量可以被加上去、也就沒有東西可以被解回來。

「直接輸出 xt1x_{t-1}」也不是好選擇,原因有二。第一,x0x_0-parametrization 把已知的結構(posterior)交給公式、只讓網路學它必須學的東西(x0x_0 長什麼樣),在每個 tt 網路面對的是同一個目標。第二,取樣時可以用同一個網路跳步:只要把 q(xt1xt,x~0)q(x_{t-1}\mid x_t,\tilde x_0) 換成 q(xtkxt,x~0)q(x_{t-k}\mid x_t,\tilde x_0)(同樣有 closed form),就能一次走 kk 步——這是下一篇 masked diffusion 幾步就取樣完的基礎。

D3PM 實務上還在 ELBO 外面加了一項輔助的 cross-entropy λE[logpθ(x0xt)]\lambda\,\mathbb E\big[-\log p_\theta(x_0\mid x_t)\big],讓網路直接對「預測 x0x_0」負責、訓練更穩。下一篇會看到:對 absorbing 鏈,ELBO 本身就會塌成這種形式的加權和,輔助項變得多餘。

先消化一下

想一想

Uniform 鏈的 Qˉt=αˉtI+(1αˉt)J\bar Q_t=\bar\alpha_tI+(1-\bar\alpha_t)J,其中 JJ 每個元素都是 1/K1/K[Qˉt]ii[\bar Q_t]_{ii}(走 tt 步後仍是原字的機率)等於:

想一想

對 absorbing 鏈,若 xtx_t 在某個位置不是 [MASK],posterior q(xt1xt,x0)q(x_{t-1}\mid x_t,x_0) 在該位置是:

想一想

x0x_0-parametrization 寫成 pθ(xt1xt)=x~0q(xt1xt,x~0)pθ(x~0xt)p_\theta(x_{t-1}\mid x_t)=\sum_{\tilde x_0}q(x_{t-1}\mid x_t,\tilde x_0)\,p_\theta(\tilde x_0\mid x_t)。這裡需要學的部分是:

想一想

離散版 ELBO 的中間項 Lt1L_{t-1} 是兩個類別分佈的 KL。它與 DDPM 對應項的差別是:

參考文獻

  1. Austin, J., Johnson, D. D., Ho, J., Tarlow, D., van den Berg, R. Structured Denoising Diffusion Models in Discrete State-Spaces. NeurIPS 2021.(D3PM:transition matrix 的一般框架、uniform / absorbing 鏈、x0x_0-parametrization、輔助 loss LλL_\lambda。)
  2. Sohl-Dickstein, J., Weiss, E. A., Maheswaranathan, N., Ganguli, S. Deep Unsupervised Learning using Nonequilibrium Thermodynamics. ICML 2015.(最早的 diffusion 論文之一,同時處理了二元狀態的離散版本。)
  3. Hoogeboom, E., Nielsen, D., Jaini, P., Forré, P., Welling, M. Argmax Flows and Multinomial Diffusion: Learning Categorical Distributions. NeurIPS 2021.
  4. Ho, J., Jain, A., Abbeel, P. Denoising Diffusion Probabilistic Models. NeurIPS 2020.(ELBO 的三種項與 x0x_0-prediction 的連續版。)