本篇重用M4.0明天的天氣只看今天:Markov Chain 與 Transition Matrix·M4.1進去就出不來:Absorbing State 與吸收時間·M1.3快篩陽性:Bayes 與 Posterior·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy
上一篇挑好了兩條鏈。
那「q(xt∣x0) 有 closed form」在離散世界是什麼意思?
先只看一個 token
一句話有很多 token,但 forward process 對每個位置是獨立地加噪聲,所以先只盯著一個位置。這個位置的狀態 x∈{1,…,K},寫成 one-hot 的列向量(row vector)x∈{0,1}K。
一條 Markov chain 每一步的規則是一個 K×K 的 transition matrix Qt:
[Qt]ij=q(xt=j∣xt−1=i),
第 i 列是「從狀態 i 出發,下一步落在各狀態的機率」,每列和為 1。用 one-hot 列向量寫,q(xt∣xt−1)=Cat(xt−1Qt)——xt−1 挑出 Qt 的那一列,那一列就是 xt 的分佈。
走 t 步就是連乘:
q(xt∣x0)=Cat(x0Qˉt),Qˉt=Q1Q2⋯Qt.
這一行是整個單元的地基。U1.1 說「forward 是我們自己定的,所以任何 t 的 xt 可以一步抽出來」——在離散世界裡,這句話就是「Qˉt 有 closed form」。一般的矩陣連乘沒有 closed form,但我們選的兩條鏈剛好都有。
這兩條鏈,Qˉt 真的算得出來
Uniform。 每步以機率 βt 把 token 換成均勻隨機的一個:
Qt=(1−βt)I+βtK111⊤.
令 J=K111⊤(每個元素都是 1/K 的矩陣),它滿足 J2=J。形如 aI+bJ 且 a+b=1 的矩陣相乘還是這個形狀:(aI+bJ)(cI+dJ)=acI+(1−ac)J。於是
Qˉt=αˉtI+(1−αˉt)J,αˉt=s≤t∏(1−βs).
讀法:到了第 t 步,token 以機率 αˉt 從沒被動過、還是原字;以機率 1−αˉt 已經被換成一個均勻隨機的字。αˉt→0 時 Qˉt→J,每個位置獨立均勻——這就是 uniform 鏈的終點。
Absorbing。 字典加一個第 K+1 個狀態 [MASK],記它的 one-hot 為 em。每步以機率 βt 把 token 變成 [MASK],[MASK] 自己不動:
Qt=(1−βt)I+βt1em⊤.
(檢查第 m 列:(1−βt)em⊤+βtem⊤=em⊤,[MASK] 一定留在 [MASK]。)同樣的代數,因為 em⊤1=1:
Qˉt=αˉtI+(1−αˉt)1em⊤.
讀法:到了第 t 步,token 以機率 αˉt 還是原字,以機率 1−αˉt 已經是 [MASK]。沒有第三種可能。 這就是上一篇說的「沒被遮的位置一定是原字」,現在它是一個公式。
兩條鏈的 αˉt 扮演的角色,和 U1.1 裡 αˉt 對 x0 的係數一樣:還剩多少原始訊號。這是我們刻意沿用同一個符號的原因。
符號這裡的 alpha bar 和連續世界的差一個平方根
連續世界寫 xt=αˉtx0+σtϵ,係數是 αˉt;這裡寫 Qˉt=αˉtI+(1−αˉt)(⋅),係數是 αˉt 本身。
差別來自「係數」的意思不同:連續世界那個是振幅(能量是它的平方),這裡是機率(token 沒被動過的機率)。兩邊都用 αˉt 是為了讓「還剩多少原始訊號」這個直覺可以直接搬過來,但代進式子的時候不要記錯。
互動 demo:Qˉt 的熱圖隨 t 演化。 拖 t:uniform 的對角線一路淡下去、其他格一起亮起來,最後整張變成一片均勻;absorbing 的質量整批流進 [MASK] 那一欄,而最後一列從頭到尾都是 [MASK] 自己。切到「Qt(只走一步)」會看到單步矩陣幾乎就是 identity——是連乘 t 次才把質量搬走的。另外對照兩條鏈的對角線:uniform 的「還是原字」比 αˉt 多一點(多了「被換掉但剛好換回自己」)。
Posterior:多知道 x0,一切都好算
U1.2 有一個關鍵步驟:反向那一步 q(xt−1∣xt) 不好算,但多給 x0 之後 q(xt−1∣xt,x0) 是一個有 closed form 的 Gaussian。同一件事在這裡更簡單,因為一切都是有限個數字。Bayes 定理:
q(xt−1∣xt,x0)∝q(xt∣xt−1)q(xt−1∣x0)=(xtQt⊤)⊙(x0Qˉt−1),
⊙ 是逐元素相乘,兩個都是長度 K 的向量(分別是「哪些 xt−1 走一步能到 xt」和「從 x0 走 t−1 步會到哪些 xt−1」),乘起來再正規化,就是 xt−1 的類別分佈。
對 absorbing 鏈這個 posterior 特別直白:若 xt 沒被遮,那 xt−1 一定也是同一個原字(機率 1);若 xt= [MASK],那 xt−1 以某個機率是 [MASK]、否則就是 x0。已知 x0 之後,反向每一步都只是「要不要把這個 [MASK] 換回 x0」的一個硬幣。
那 loss 長什麼樣?
現在寫訓練目標。U1.2 為了盡快到達那五行程式碼,直接用「預測 x0 的 MSE」當 loss,把 DDPM 原本的 variational bound 放在一旁。在離散世界裡值得把這個 bound 攤開看一次,因為它的每一項都是有限個數字之間的 KL,沒有任何近似。
Reverse model 是 pθ(xt−1∣xt),起點 p(xT) 是鏈的終點分佈。DDPM [4] 的 bound 一字不改地成立 [1, 2]:
−logpθ(x0)≤L0Eq[−logpθ(x0∣x1)]+t=2∑TLt−1EqKL(q(xt−1∣xt,x0)pθ(xt−1∣xt))+LTKL(q(xT∣x0)p(xT)).
LT 沒有參數,而且對我們選的兩條鏈它趨近 0(αˉT≈0 時 q(xT∣x0) 就是終點分佈)。L0 是一個普通的 cross-entropy。中間每一項 Lt−1 是兩個 K 維類別分佈的 KL——左邊是那條 Bayes 算出來的 posterior q(xt−1∣xt,x0),右邊是網路。
展開細節推導的骨架:為什麼 bound 長這樣(與 DDPM 逐字相同)
從 logpθ(x0)=log∫pθ(x0:T)dx1:T 出發,用 q(x1:T∣x0) 當 proposal 套 Jensen 不等式:
−logpθ(x0)≤Eq(x1:T∣x0)[logpθ(x0:T)q(x1:T∣x0)].把 q(x1:T∣x0)=∏tq(xt∣xt−1) 用 Bayes 改寫成 q(xT∣x0)∏t≥2q(xt−1∣xt,x0),pθ(x0:T)=p(xT)∏tpθ(xt−1∣xt),逐項配對,就得到上面三種項。在連續世界裡積分換成離散世界的求和,其餘一步都不用改;這就是為什麼說「bound 一字不改地成立」。
網路該輸出什麼
還有一個設計決定:pθ(xt−1∣xt) 要怎麼參數化?
最直接的想法是讓網路直接輸出 xt−1 的 K 個 logits。D3PM [1] 的選擇不同,而且是這一篇最值得記的一句話:讓網路預測 x0,再用已知的 posterior 把它折回 xt−1。
pθ(xt−1∣xt)=x~0∑q(xt−1∣xt,x~0)pθ(x~0∣xt).
網路輸出的是 pθ(x~0∣xt) 的 logits(每個位置 K 個數);q(xt−1∣xt,x~0) 是上面那個 closed form,不用學。這叫 x0-parametrization,它就是 U1.2 裡 x0-prediction 的離散翻譯——在那裡網路預測 x^0、再用 q(xt−1∣xt,x^0) 的 Gaussian posterior 往回走一步;在這裡完全一樣,只是「預測 x0」從輸出一個向量變成輸出一個類別分佈。
課堂提問Q1
U1.2 說 x0-、ϵ-、v-prediction 是三個線性可互換的座標,實務上常選 ϵ。離散世界裡為什麼只剩 x0-parametrization 這一種選擇?
先想一想,再展開看整理後的答案
因為那三個座標可以互換,靠的是 xt=αˉtx0+σtϵ 這條 affine 關係——知道 xt 與任一個,就能解出另外兩個。離散世界沒有這條關係:xt 是從 x0 抽出來的一個類別,沒有一個叫「ϵ」的連續量可以被加上去、也就沒有東西可以被解回來。
「直接輸出 xt−1」也不是好選擇,原因有二。第一,x0-parametrization 把已知的結構(posterior)交給公式、只讓網路學它必須學的東西(x0 長什麼樣),在每個 t 網路面對的是同一個目標。第二,取樣時可以用同一個網路跳步:只要把 q(xt−1∣xt,x~0) 換成 q(xt−k∣xt,x~0)(同樣有 closed form),就能一次走 k 步——這是下一篇 masked diffusion 幾步就取樣完的基礎。
D3PM 實務上還在 ELBO 外面加了一項輔助的 cross-entropy λE[−logpθ(x0∣xt)],讓網路直接對「預測 x0」負責、訓練更穩。下一篇會看到:對 absorbing 鏈,ELBO 本身就會塌成這種形式的加權和,輔助項變得多餘。
先消化一下
參考文獻
- 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 鏈、x0-parametrization、輔助 loss Lλ。)
- Sohl-Dickstein, J., Weiss, E. A., Maheswaranathan, N., Ganguli, S. Deep Unsupervised Learning using Nonequilibrium Thermodynamics. ICML 2015.(最早的 diffusion 論文之一,同時處理了二元狀態的離散版本。)
- Hoogeboom, E., Nielsen, D., Jaini, P., Forré, P., Welling, M. Argmax Flows and Multinomial Diffusion: Learning Categorical Distributions. NeurIPS 2021.
- Ho, J., Jain, A., Abbeel, P. Denoising Diffusion Probabilistic Models. NeurIPS 2020.(ELBO 的三種項與 x0-prediction 的連續版。)