U4.3 13 分鐘閱讀 2026年9月

U4.3 因子化誤差:離散世界的曲率

本篇重用M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy·M1.1從鞋子猜身高:Conditional Expectation

上一篇說步數可以壓到 8 步、甚至 1 步。
那同時翻開很多格,代價是什麼?

網路給的是什麼

上一篇結尾留了一個問題:步數越少、每步同時翻開的格子越多,會出什麼問題?先把「網路給的是什麼」說精確。

網路看到 xtx_t,對每個被遮的位置 \ell 輸出一個分佈 pθ(x0xt)p_\theta(x_0^\ell\mid x_t)。這是一個 marginal——「在看到 xtx_t 的前提下,位置 \ell 單獨來看是哪個字的機率」。它沒有告訴我們位置 \ell 和位置 \ell' 一起會是什麼;網路沒有輸出任何 joint。

取樣時如果這一步只翻開一個位置,marginal 就是我們需要的全部。但如果同時翻開一組位置 SS,我們實際上是從

Spθ(x0xt)\prod_{\ell\in S}p_\theta(x_0^\ell\mid x_t)

抽樣。而正確的目標是 conditional joint p(x0Sxt)p(x_0^S\mid x_t)。兩者相等的條件是:給定 xtx_tSS 裡的位置彼此 conditionally independent。 語言裡這幾乎從不成立——「今天想吃 [MASK] [MASK]」的兩個空格,「牛肉/麵」和「壽/司」各自都合理,「牛肉/司」不合理,而 marginal 的乘積會以不小的機率生出它。

這就是同時翻開多個 token 的全部代價:用乘積代替了 joint,忽略了條件相依。

課堂提問Q1

既然如此,最省事的做法——一步把所有 token 全部翻開——會壞到什麼程度?請先在下面這個小例子上算:長度 8 的 0/1 序列,資料分佈是所有偶數個 1 的序列(共 128 條,均勻)。假設網路是完美的(輸出的 marginal 就是真實的 conditional marginal),從全 [MASK] 出發一步翻開全部,得到合法序列的機率是多少?如果改成每步翻開一個、走 8 步呢?每步翻開兩個、走 4 步呢?

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

從全 [MASK] 出發,每個位置的真實 marginal 都是 12\tfrac12(偶數 parity 的序列裡每個位置是 0 或 1 的機率各半)。一步翻開全部,就是從 {0,1}8\{0,1\}^8 均勻抽一條,落在偶數 parity 的機率是 12\tfrac12。一半的樣本不合法。

每步翻開一個、走 8 步:前 7 個位置的 marginal 依然各是 12\tfrac12(少於 8 個已知的位置時,剩下的位置仍可以配成任一 parity),第 8 個位置的 conditional 是確定的——它必須把 parity 補成偶數。完美的網路會輸出一個機率 1 的 marginal,合法率 1

每步翻開兩個、走 4 步:前三步都沒問題(任何少於 8 個位置的子集,joint 就是均勻的乘積),但最後一步同時翻開最後兩個位置——它們的 conditional joint 只允許兩種組合(01 或 10,或 00 與 11,取決於前六個的 parity),marginal 的乘積卻允許四種。合法率 12\tfrac12

把三個數字放在一起會看到這個例子想說的事:誤差不是由步數均勻地決定,它出現在「同時翻開的位置之間存在條件相依」的那一步。parity 的相依全部藏在最後一個自由度上,所以只有最後一步翻開超過一個位置時才會出錯;真實語言的相依散在各處(相鄰字、主謂一致、括號配對),所以每一步都會漏一點。步數越少、每步翻開越多,漏掉的相依越多。

把誤差寫成一個數

上面的例子可以寫成一般式。在某一步,模型從 xtx_t 同時翻開位置集合 SS。這一步應該抽的是 p(x0Sxt)p(x_0^S\mid x_t),實際抽的是 Sp(x0xt)\prod_{\ell\in S}p(x_0^\ell\mid x_t)(先假設網路完美,只看因子化本身的代價)。兩者之間的差距:

FE(Sxt)  =  KL(p(x0Sxt)    Sp(x0xt)).\mathrm{FE}(S\mid x_t)\;=\;\mathrm{KL}\Big(p(x_0^S\mid x_t)\;\Big\|\;\prod_{\ell\in S}p(x_0^\ell\mid x_t)\Big).

這個量在資訊理論裡叫 total correlation(多變數版本的 mutual information):它衡量一組變數「合起來看」比「分開看」多知道多少。它 0\ge0,且等於零若且唯若 SS 裡的位置在給定 xtx_t 下 conditionally independentS=1|S|=1 時它恆為零——一次翻開一個位置永遠沒有因子化誤差。

整個取樣程序的因子化誤差,是每一步這個量的和(對每步的 xtx_t 取期望)。parity 例子裡,前面各步的 FE\mathrm{FE} 都是零,最後一步翻開兩個位置時 FE=log2\mathrm{FE}=\log 2——正好對應「一半樣本不合法」。

展開細節為什麼是 KL、以及它和 loss 的關係

取樣程序真正生成的分佈是把每一步的乘積分佈串起來得到的 p~(x0)\tilde p(x_0)。用 chain rule 把 logp(x0)p~(x0)\log\frac{p(x_0)}{\tilde p(x_0)} 沿著翻開的順序拆開,每一步貢獻的正是 logp(x0Sxt)p(x0xt)\log\frac{p(x_0^{S}\mid x_t)}{\prod_\ell p(x_0^\ell\mid x_t)},取期望就是上面的 KL。所以 KL(pp~)=stepsEFE(Sxt)\mathrm{KL}(p\,\|\,\tilde p)=\sum_{\text{steps}}\mathbb E\,\mathrm{FE}(S\mid x_t)——生成分佈與真實分佈的差距,就是所有步的因子化誤差之和(網路完美時)。

值得注意的是這個誤差不在訓練 loss 裡U4.2 的 ELBO 是對「每步翻開一個位置」的鏈算的(TT\to\infty 時每步幾乎只翻開零或一個位置),它的最小值是完美的 marginal。因子化誤差純粹是取樣時「一步走太大」造成的——這和 U3.2 的情形一模一樣:訓練目標是精確的速度場,Euler 誤差是離散化造成的。

互動 demo:marginal 的乘積 vs 真實的 joint。 就是上面那個 parity 例子,只是把兩張表並排畫出來。按 k=1k=1:每一步只翻開一格,FE\mathrm{FE} 一直是 0,跑 1000 次的合法率 100%。按 k=2k=2 或更大:FE\mathrm{FE} 在前面幾步還是 0,只有最後那一步跳到 log20.693\log 2\approx0.693,合法率掉到 50% 上下(實測 k=2k=2 是 49.8%、k=4k=4 是 48.3%、k=8k=8 是 52.6%)。k=8k=8 時兩張表會畫成一排細柱:真實 joint 的 256 個組合裡只有 128 個機率不是 0,marginal 的乘積卻 256 個都有值

這是離散世界的曲率

把這一篇和 U3.2 並排,會看到同一個結構:

連續(U3離散(本單元)
訓練學到的物件精確的速度場 ut(x)u_t(x)精確的 marginal p(x0xt)p(x_0^\ell\mid x_t)
一步走太大時假設了什麼軌跡在這一步內是直線這一步翻開的位置條件獨立
誤差的來源曲率 x¨tdt\int\|\ddot x_t\|\,dt條件相依(total correlation)
何時為零軌跡是直線每步只翻開一個位置
減少誤差的做法多走幾步;拉直軌跡(reflow、OT 配對);高階 solver多走幾步;挑「相依弱」的位置先翻(信心排序);顯式建模相依

這張表是這個單元最重要的一張。它說的是:有限步取樣的誤差,在兩個世界裡都來自「一步之內假設了太簡單的結構」。 連續世界假設直線,離散世界假設獨立。U3.3U3.4 在連續世界做的「拉直」,在離散世界沒有直接對應——沒有「直的配對」這回事——所以離散世界減少誤差的手段目前主要落在最後一行的後兩項:MaskGIT [3] 用信心排序挑位置;近期的工作(例如 Discrete Copula Diffusion [4])則試著在取樣時補回相依的結構。這條線還在發展中。

那不如直接用 autoregressive?

表格最後一列還說了另一件事。「每步只翻開一個位置」在因子化誤差上是零,而它就是 autoregressive(AR)模型的做法——只是 AR 固定從左到右,masked diffusion 每步隨機挑一個位置。Ou et al. [1] 與更早的 ARDM [2] 都指出:absorbing diffusion 學到的 p(x0xt)p(x_0^\ell\mid x_t) 是乾淨資料的所有 conditional,所以每步翻一個的 masked diffusion 就是一個任意順序的 AR 模型。

於是兩者的取捨可以講得很乾脆:

  • AR:一次一個 token,固定順序,factorization 精確、沒有因子化誤差;代價是 LL 個 token 要 LL 次網路呼叫,無法平行
  • Masked diffusion:任意順序,一步可以翻開很多個、可以平行L/kL/k 次呼叫;代價是每步翻開的位置之間的因子化誤差。此外它天生能做填空(雙向 context),AR 需要額外設計。

這是 U3.2 那個「步數 vs 誤差」的取捨換了一套詞。Zheng et al. [5] 提醒了一件實務上的事:比較 masked diffusion 與 AR 的 perplexity 時,若取樣步數不夠、或類別抽樣的數值實作不精確,會把因子化誤差誤讀成別的東西;比較時要把「每步翻開幾個」當成明確的實驗變數。

先消化一下

想一想

Masked diffusion 一步同時翻開位置集合 SS,等價於假設:

想一想

Parity toy 裡,每步翻開兩個、走 4 步的合法率是 12\tfrac12,而不是比一步全翻開好。這說明:

想一想

連續世界的有限步誤差來自曲率,離散世界的來自因子化。兩者共同的結構是:

想一想

「每步只翻開一個位置」的 masked diffusion 與 autoregressive 模型的關係是:

參考文獻

  1. 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.(absorbing diffusion 學到的是乾淨資料的 conditional;與任意順序 AR 的關係。)
  2. Hoogeboom, E., Gritsenko, A. A., Bastings, J., Poole, B., van den Berg, R., Salimans, T. Autoregressive Diffusion Models. ICLR 2022.(ARDM:任意順序 AR 與 absorbing diffusion 的等價。)
  3. Chang, H., Zhang, H., Jiang, L., Liu, C., Freeman, W. T. MaskGIT: Masked Generative Image Transformer. CVPR 2022.(信心排序的平行解碼。)
  4. Liu, A., Liu, O., Van den Broeck, G. Discrete Copula Diffusion. ICLR 2025.(把平行解碼漏掉的相依顯式補回。)
  5. Zheng, K., Chen, Y., Mao, H., Liu, M.-Y., Zhu, J., Zhang, Q. Masked Diffusion Models are Secretly Time-Agnostic Masked Models and Exploit Inaccurate Categorical Sampling. ICLR 2025.(比較 MDM 與 AR 時的實驗陷阱。)
  6. Sahoo, S. S. et al. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.