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

U4.5 實作:Parity 與 Markov Toy 上的 Masked Diffusion

本篇重用M4.0明天的天氣只看今天:Markov Chain 與 Transition Matrix·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy

為什麼是這兩個 toy

U1.5 選 2D 月牙,是因為真實的 ptp_t、score 可以精確算。離散世界的對應要求是:狀態空間小到可以把所有 conditional marginal p(x0xt)p(x_0^\ell\mid x_t) 用查表算出來,這樣網路學得好不好有標準答案,而且因子化誤差可以和模型誤差分開量。

  • Parity toy:長度 L=8L=8 的 0/1 序列,資料是全部 128 條偶數 parity 的序列(均勻)。相依全部藏在「最後一個自由度」上,U4.3 Q1 的三個數字(12\tfrac121112\tfrac12)可以直接驗證。
  • Markov toy:長度 L=16L=16 的 0/1 序列,第一格均勻,之後每格以機率 0.90.9 等於上一格。相依在相鄰位置之間、處處都有,「步數 vs 誤差」的曲線是漸變的。狀態空間 2162^{16},查表仍可行。

網路:小型 transformer(2 層、寬 64、雙向 attention),輸入 xtx_t(含 [MASK] 符號的 K+1=3K+1=3 類 embedding)與 tt 的 embedding,對每個位置輸出 K=2K=2 個 logits。若採用 U4.2 說的「網路不需要看 tt」,可以把 tt 拿掉當作業 4 的對照。T=32T=32,線性 schedule αˉt=1t/T\bar\alpha_t=1-t/T

步驟 1:forward 與訓練

MASK = 2                                      # 第 K+1 個符號

def q_sample(x0, t):                          # absorbing forward 的 closed form
    keep = torch.rand_like(x0.float()) < alpha_bar[t][:, None]
    return torch.where(keep, x0, torch.full_like(x0, MASK))

def weight(t):                                # (ᾱ_{t-1} − ᾱ_t) / (1 − ᾱ_t)
    return (alpha_bar[t - 1] - alpha_bar[t]) / (1 - alpha_bar[t])

for step in range(n_steps):                   # 加權 masked cross-entropy
    x0 = sample_data(batch)
    t  = torch.randint(1, T + 1, (batch,))
    xt = q_sample(x0, t)
    logits = model(xt, t)                     # (batch, L, K)
    ce = F.cross_entropy(logits.transpose(1, 2), x0, reduction='none')  # (batch, L)
    masked = (xt == MASK).float()
    loss = (weight(t)[:, None] * ce * masked).sum(1).mean()
    loss.backward(); opt.step(); opt.zero_grad()

U1.5 的五行逐行對照:q_sample 從「乘 αˉ\sqrt{\bar\alpha} 加噪聲」變成「以 αˉ\bar\alpha 的機率保留、否則遮掉」;loss 從 MSE 變成只在被遮位置計算、帶時間權重的 cross-entropy。其他一行都沒變。

檢查點:Markov toy 上,把 logits.softmax(-1)exact_conditional(xt) 在被遮位置比對,平均 KL 應在幾千步內降到 10210^{-2} 以下。parity toy 上,注意「只剩一個被遮位置」的樣本:網路應學到接近確定的輸出(其他情形下 marginal 都該是 0.5)。

步驟 2:模型誤差 vs 因子化誤差(圖 a)

這一步要把兩種誤差分開。對 Markov toy:

  • 模型誤差:對一批 xtx_t,算 KL(exactmodel)\mathrm{KL}\big(\text{exact}\,\|\,\text{model}\big) 在被遮位置的平均,對 tt 畫曲線。期待:tt 大(幾乎全遮)時網路容易——marginal 接近 0.5;tt 中段最難。
  • 因子化誤差:用 exact_conditional不用網路)跑取樣,每步翻開 kk 個隨機位置,對 k{1,2,4,8,16}k\in\{1,2,4,8,16\}kl_to_data。這條曲線只含因子化誤差。

把兩條線放在同一張圖上:模型誤差幾乎不隨 kk 變,因子化誤差隨 kk 單調上升。這張圖是 U4.3 那張對照表的實證。

一張圖上兩種誤差,
怎麼一眼看出哪一段是誰的責任?

互動 demo:兩種誤差分開看。 這是圖 a 的「標準答案版」:Markov toy、conditional marginal 用前向後向精確算,模型誤差用一個旋鈕 ϵ\epsilon 模擬(把 marginal 往均勻拉 ϵ\epsilon)。縱軸用一個好懂的統計量——樣本的切換率(相鄰兩格不同的比例),資料的真值是 0.10。讀法就兩句話:三條線在 k=1k=1 的高低差是模型誤差(完美網路實測 0.09~0.11,ϵ=0.15\epsilon=0.15 是 0.17~0.18),同一條線隨 kk 爬升的那一段是因子化誤差(完美網路一路爬到 k=16k=160.50,也就是「每格獨立擲硬幣」)。

補充為什麼用切換率,而不是直接算 KL

Markov toy 的狀態空間是 2162^{16},理論上可以查表算精確的 KL,但用有限樣本估 KL 會有很大的偏差(沒出現過的序列該算多少?),在 kk 小的時候那個偏差會蓋掉真正的訊號。

切換率是這個資料分佈的一個充分好用的摘要:它有一個已知的真值(0.10),而因子化誤差恰好會讓它往「每格獨立」的 0.50 移動,所以偏差的方向和大小都讀得出來。實作時建議兩個都算:KL 用來確認整體,切換率用來看趨勢。

步驟 3:每步翻開幾個 vs 合法率(圖 b)

@torch.no_grad()
def sample(model, n, steps, chain='absorbing'):
    xt = torch.full((n, L), MASK)
    ts = torch.linspace(T, 0, steps + 1).long()      # 跳步:T → … → 0
    for t, s in zip(ts[:-1], ts[1:]):
        probs = model(xt, t.expand(n)).softmax(-1)
        flip  = torch.rand(n, L) < (alpha_bar[s] - alpha_bar[t]) / (1 - alpha_bar[t])
        new   = torch.distributions.Categorical(probs).sample()
        xt = torch.where((xt == MASK) & flip, new, xt)   # 翻開就固定
    return xt

對 parity toy,步數 {1,2,4,8}\{1,2,4,8\} 各抽 2000 條,畫 legal 的比例。期待:1 步是 50%(U4.3 Q1 算過的數字),2 步幾乎沒有進步(約 52%),4 步約 64%,8 步爬到 78–79%——但不是 100%。8 步為什麼不是 100%、而且為什麼是卡在 79% 而不是別的數,正是下面 Q1 要問的事,先寫下你的猜測再往下讀。倒過來,如果連 1、2 步都明顯低於 50%,那才是網路本身有問題,回步驟 1 的檢查點。

對 Markov toy,步數 {1,2,4,8,16}\{1,2,4,8,16\},畫 kl_to_data。期待一條隨步數下降、漸變的曲線。把這張圖留好,U5 加上 remasking 之後會在同一張圖上疊新的線。

步驟 4:absorbing vs uniform 的修正實驗(圖 c)

另訓一個 uniform 鏈的模型(q_sample 改成「以 1αˉt1-\bar\alpha_t 的機率換成均勻隨機字」,loss 改回 U4.1 的一般 KL 形式——參考解答提供 d3pm_loss)。在 Markov toy 上重做 U4.4 demo 的實驗:第 2 步強制植入一個與鄰格相反的字,之後照常取樣,統計 200 次裡錯字最後被改掉的比例,以及最終樣本的 kl_to_data

期待:absorbing 的修正率是 0%;uniform 有非零的修正率,但整體 kl_to_data 未必更好——把兩個數字並排,寫兩句話說明 U4.4 講的「修正能力 vs 乾淨 context」在這個 toy 上是哪一邊贏。

課堂提問Q1

步驟 3 在 parity toy 上用 8 步、L=8L=8,所以「每步平均翻開一格」。但實際跑出來的合法率是 78–79%,不是 100%。

哪裡出問題?

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

不是模型,也不是實作 bug。問題在「平均翻開一格」和「每步剛好翻開一格」不是同一件事。

看那一行程式:flip = torch.rand(n, L) < (ᾱ_s − ᾱ_t)/(1 − ᾱ_t)。每個位置各自擲一枚硬幣,所以一步之內翻開幾格是一個 Binomial 隨機變數——期望值是 1,但有相當的機率是 0 或 2 以上

parity 的相依全部藏在最後一個自由度,所以真正該問的不是「有沒有哪一步翻開兩格」,而是最後一次翻開了幾格。只要還有兩格以上被遮,每一格的 marginal 就恰好是 0.5——網路學得再好也只能給 0.5,這是 parity 這個資料本身的性質;只有「剩一格」時,那一格才被 parity 決定。

這條樣本合法,等價於最後一次翻開只翻了一格。

否則最後那一批全是公平銅板,parity 有一半的機會湊成偶數。

這個 toy 小到可以把數字算死。線性 schedule 之下第 ii 步的翻開機率是 1Ni\frac{1}{N-i},而這一組機率等價於每個位置各自獨立抽一個 {0,,N1}\{0,\dots,N-1\} 上的均勻標籤——位置 \ell 第一次被翻開在第 ii 步的機率是 j<iN1jNj1Ni=1N\prod_{j<i}\frac{N-1-j}{N-j}\cdot\frac{1}{N-i}=\frac1N,和 ii 無關。於是「最後一次只翻一格」就是「最大的那個標籤只有一個位置拿到」:

合法率=1+q2,q=Pr[最大標籤唯一]=LNLk=0N1kL1.\text{合法率}=\frac{1+q}{2},\qquad q=\Pr[\text{最大標籤唯一}]=\frac{L}{N^{L}}\sum_{k=0}^{N-1}k^{\,L-1}.

L=N=8L=N=8 代進去,q=0.5723q=0.5723、合法率 78.6%。所以量到 78–79% 不是實作沒調好——那就是這個設定的天花板,剩下的 21.4% 全部是因子化誤差。

三種修法,各對應不同的意思:

  1. 強制每步剛好一格(在還被遮的位置裡隨機挑一個)。這樣因子化誤差嚴格是 0,合法率會是 100%——這就是 U4.3 的 demo 做的事,也等於一個任意順序的 autoregressive 模型。
  2. 步數加大(例如 32 步)。翻開機率變小、最後一次只翻一格的機會上升,合法率往 100% 靠:同一條公式給出 N=16N=16 是 88.4%、N=32N=32 是 94.0%、N=64N=64 是 96.9%。這是「多走幾步」的標準答案,而且收斂得比想像中慢。
  3. 什麼都不改,但把它記在報告裡。 這才是這一題真正的用意:「步數 =L=L」不等於「每步一格」,而 masked diffusion 的論文報數字時如果只寫步數、沒寫每步實際翻開幾格,那個數字就不能直接和 autoregressive 比。U4.3 引的 Zheng et al. 提醒的正是這件事。

作業

  1. Uniform vs absorbing 的合法率。在 parity toy 上對兩條鏈各畫步數 vs 合法率。uniform 鏈在 8 步時是否也接近 100%?若不是,說明是哪一種誤差(提示:uniform 鏈的 xtx_t 裡「只剩一個不確定的位置」這個狀態,模型認得出來嗎?)。
  2. 信心排序。把步驟 3 的 flip 改成「每步翻開 marginal 最尖(entropy 最低)的 kk 個位置」(MaskGIT 式),重畫 Markov toy 的曲線。改善多少?在 parity toy 上呢——為什麼那裡幾乎沒有幫助?
  3. 換 parametrization。把網路改成直接預測 xt1x_{t-1}(而非 x0x_0),用 U4.1 的一般 ELBO 訓練。比較:(a) 訓練是否更難收斂;(b) 能不能跳步取樣。寫下你觀察到的差別與 U4.1 Q1 的對應。
  4. (選)拿掉 tt。把網路的 tt 輸入移除、重訓 absorbing 模型。loss 曲線與步驟 3 的合法率是否改變?用 U4.2 那個 <Details> 裡的論證解釋為什麼。
想一想

步驟 1 的 loss 只在被遮位置計算。若把沒被遮的位置也算進 cross-entropy,會發生什麼?

想一想

Parity toy 上 8 步取樣的合法率是 78–79%,而不是接近 100%。最可能的原因是:

參考文獻

  1. Sahoo, S. S. et al. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.(步驟 1 的 loss 與步驟 3 的取樣程序。)
  2. Austin, J. et al. Structured Denoising Diffusion Models in Discrete State-Spaces. NeurIPS 2021.(步驟 4 uniform 鏈的一般 loss。)