U4.3 因子化誤差:離散世界的曲率
本篇重用M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy·M1.1從鞋子猜身高:Conditional Expectation
上一篇說步數可以壓到 8 步、甚至 1 步。
那同時翻開很多格,代價是什麼?
網路給的是什麼
上一篇結尾留了一個問題:步數越少、每步同時翻開的格子越多,會出什麼問題?先把「網路給的是什麼」說精確。
網路看到 ,對每個被遮的位置 輸出一個分佈 。這是一個 marginal——「在看到 的前提下,位置 單獨來看是哪個字的機率」。它沒有告訴我們位置 和位置 一起會是什麼;網路沒有輸出任何 joint。
取樣時如果這一步只翻開一個位置,marginal 就是我們需要的全部。但如果同時翻開一組位置 ,我們實際上是從
抽樣。而正確的目標是 conditional joint 。兩者相等的條件是:給定 , 裡的位置彼此 conditionally independent。 語言裡這幾乎從不成立——「今天想吃 [MASK] [MASK]」的兩個空格,「牛肉/麵」和「壽/司」各自都合理,「牛肉/司」不合理,而 marginal 的乘積會以不小的機率生出它。
這就是同時翻開多個 token 的全部代價:用乘積代替了 joint,忽略了條件相依。
課堂提問Q1
既然如此,最省事的做法——一步把所有 token 全部翻開——會壞到什麼程度?請先在下面這個小例子上算:長度 8 的 0/1 序列,資料分佈是所有偶數個 1 的序列(共 128 條,均勻)。假設網路是完美的(輸出的 marginal 就是真實的 conditional marginal),從全 [MASK] 出發一步翻開全部,得到合法序列的機率是多少?如果改成每步翻開一個、走 8 步呢?每步翻開兩個、走 4 步呢?
先想一想,再展開看整理後的答案
從全 [MASK] 出發,每個位置的真實 marginal 都是 (偶數 parity 的序列裡每個位置是 0 或 1 的機率各半)。一步翻開全部,就是從 均勻抽一條,落在偶數 parity 的機率是 。一半的樣本不合法。
每步翻開一個、走 8 步:前 7 個位置的 marginal 依然各是 (少於 8 個已知的位置時,剩下的位置仍可以配成任一 parity),第 8 個位置的 conditional 是確定的——它必須把 parity 補成偶數。完美的網路會輸出一個機率 1 的 marginal,合法率 1。
每步翻開兩個、走 4 步:前三步都沒問題(任何少於 8 個位置的子集,joint 就是均勻的乘積),但最後一步同時翻開最後兩個位置——它們的 conditional joint 只允許兩種組合(01 或 10,或 00 與 11,取決於前六個的 parity),marginal 的乘積卻允許四種。合法率 。
把三個數字放在一起會看到這個例子想說的事:誤差不是由步數均勻地決定,它出現在「同時翻開的位置之間存在條件相依」的那一步。parity 的相依全部藏在最後一個自由度上,所以只有最後一步翻開超過一個位置時才會出錯;真實語言的相依散在各處(相鄰字、主謂一致、括號配對),所以每一步都會漏一點。步數越少、每步翻開越多,漏掉的相依越多。
把誤差寫成一個數
上面的例子可以寫成一般式。在某一步,模型從 同時翻開位置集合 。這一步應該抽的是 ,實際抽的是 (先假設網路完美,只看因子化本身的代價)。兩者之間的差距:
這個量在資訊理論裡叫 total correlation(多變數版本的 mutual information):它衡量一組變數「合起來看」比「分開看」多知道多少。它 ,且等於零若且唯若 裡的位置在給定 下 conditionally independent。 時它恆為零——一次翻開一個位置永遠沒有因子化誤差。
整個取樣程序的因子化誤差,是每一步這個量的和(對每步的 取期望)。parity 例子裡,前面各步的 都是零,最後一步翻開兩個位置時 ——正好對應「一半樣本不合法」。
展開細節為什麼是 KL、以及它和 loss 的關係
互動 demo:marginal 的乘積 vs 真實的 joint。 就是上面那個 parity 例子,只是把兩張表並排畫出來。按 :每一步只翻開一格, 一直是 0,跑 1000 次的合法率 100%。按 或更大: 在前面幾步還是 0,只有最後那一步跳到 ,合法率掉到 50% 上下(實測 是 49.8%、 是 48.3%、 是 52.6%)。 時兩張表會畫成一排細柱:真實 joint 的 256 個組合裡只有 128 個機率不是 0,marginal 的乘積卻 256 個都有值。
這是離散世界的曲率
把這一篇和 U3.2 並排,會看到同一個結構:
| 連續(U3) | 離散(本單元) | |
|---|---|---|
| 訓練學到的物件 | 精確的速度場 | 精確的 marginal |
| 一步走太大時假設了什麼 | 軌跡在這一步內是直線 | 這一步翻開的位置條件獨立 |
| 誤差的來源 | 曲率 | 條件相依(total correlation) |
| 何時為零 | 軌跡是直線 | 每步只翻開一個位置 |
| 減少誤差的做法 | 多走幾步;拉直軌跡(reflow、OT 配對);高階 solver | 多走幾步;挑「相依弱」的位置先翻(信心排序);顯式建模相依 |
這張表是這個單元最重要的一張。它說的是:有限步取樣的誤差,在兩個世界裡都來自「一步之內假設了太簡單的結構」。 連續世界假設直線,離散世界假設獨立。U3.3、U3.4 在連續世界做的「拉直」,在離散世界沒有直接對應——沒有「直的配對」這回事——所以離散世界減少誤差的手段目前主要落在最後一行的後兩項:MaskGIT [3] 用信心排序挑位置;近期的工作(例如 Discrete Copula Diffusion [4])則試著在取樣時補回相依的結構。這條線還在發展中。
那不如直接用 autoregressive?
表格最後一列還說了另一件事。「每步只翻開一個位置」在因子化誤差上是零,而它就是 autoregressive(AR)模型的做法——只是 AR 固定從左到右,masked diffusion 每步隨機挑一個位置。Ou et al. [1] 與更早的 ARDM [2] 都指出:absorbing diffusion 學到的 是乾淨資料的所有 conditional,所以每步翻一個的 masked diffusion 就是一個任意順序的 AR 模型。
於是兩者的取捨可以講得很乾脆:
- AR:一次一個 token,固定順序,factorization 精確、沒有因子化誤差;代價是 個 token 要 次網路呼叫,無法平行。
- Masked diffusion:任意順序,一步可以翻開很多個、可以平行, 次呼叫;代價是每步翻開的位置之間的因子化誤差。此外它天生能做填空(雙向 context),AR 需要額外設計。
這是 U3.2 那個「步數 vs 誤差」的取捨換了一套詞。Zheng et al. [5] 提醒了一件實務上的事:比較 masked diffusion 與 AR 的 perplexity 時,若取樣步數不夠、或類別抽樣的數值實作不精確,會把因子化誤差誤讀成別的東西;比較時要把「每步翻開幾個」當成明確的實驗變數。
先消化一下
參考文獻
- 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 的關係。)
- Hoogeboom, E., Gritsenko, A. A., Bastings, J., Poole, B., van den Berg, R., Salimans, T. Autoregressive Diffusion Models. ICLR 2022.(ARDM:任意順序 AR 與 absorbing diffusion 的等價。)
- Chang, H., Zhang, H., Jiang, L., Liu, C., Freeman, W. T. MaskGIT: Masked Generative Image Transformer. CVPR 2022.(信心排序的平行解碼。)
- Liu, A., Liu, O., Van den Broeck, G. Discrete Copula Diffusion. ICLR 2025.(把平行解碼漏掉的相依顯式補回。)
- 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 時的實驗陷阱。)
- Sahoo, S. S. et al. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.