U4.5 實作:Parity 與 Markov Toy 上的 Masked Diffusion
本篇重用M4.0明天的天氣只看今天:Markov Chain 與 Transition Matrix·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy
為什麼是這兩個 toy
U1.5 選 2D 月牙,是因為真實的 、score 可以精確算。離散世界的對應要求是:狀態空間小到可以把所有 conditional marginal 用查表算出來,這樣網路學得好不好有標準答案,而且因子化誤差可以和模型誤差分開量。
- Parity toy:長度 的 0/1 序列,資料是全部 128 條偶數 parity 的序列(均勻)。相依全部藏在「最後一個自由度」上,U4.3 Q1 的三個數字(、、)可以直接驗證。
- Markov toy:長度 的 0/1 序列,第一格均勻,之後每格以機率 等於上一格。相依在相鄰位置之間、處處都有,「步數 vs 誤差」的曲線是漸變的。狀態空間 ,查表仍可行。
網路:小型 transformer(2 層、寬 64、雙向 attention),輸入 (含 [MASK] 符號的 類 embedding)與 的 embedding,對每個位置輸出 個 logits。若採用 U4.2 說的「網路不需要看 」,可以把 拿掉當作業 4 的對照。,線性 schedule 。
步驟 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 從「乘 加噪聲」變成「以 的機率保留、否則遮掉」;loss 從 MSE 變成只在被遮位置計算、帶時間權重的 cross-entropy。其他一行都沒變。
檢查點:Markov toy 上,把 logits.softmax(-1) 與 exact_conditional(xt) 在被遮位置比對,平均 KL 應在幾千步內降到 以下。parity toy 上,注意「只剩一個被遮位置」的樣本:網路應學到接近確定的輸出(其他情形下 marginal 都該是 0.5)。
步驟 2:模型誤差 vs 因子化誤差(圖 a)
這一步要把兩種誤差分開。對 Markov toy:
- 模型誤差:對一批 ,算 在被遮位置的平均,對 畫曲線。期待: 大(幾乎全遮)時網路容易——marginal 接近 0.5; 中段最難。
- 因子化誤差:用
exact_conditional(不用網路)跑取樣,每步翻開 個隨機位置,對 算kl_to_data。這條曲線只含因子化誤差。
把兩條線放在同一張圖上:模型誤差幾乎不隨 變,因子化誤差隨 單調上升。這張圖是 U4.3 那張對照表的實證。
一張圖上兩種誤差,
怎麼一眼看出哪一段是誰的責任?
互動 demo:兩種誤差分開看。 這是圖 a 的「標準答案版」:Markov toy、conditional marginal 用前向後向精確算,模型誤差用一個旋鈕 模擬(把 marginal 往均勻拉 )。縱軸用一個好懂的統計量——樣本的切換率(相鄰兩格不同的比例),資料的真值是 0.10。讀法就兩句話:三條線在 的高低差是模型誤差(完美網路實測 0.09~0.11, 是 0.17~0.18),同一條線隨 爬升的那一段是因子化誤差(完美網路一路爬到 的 0.50,也就是「每格獨立擲硬幣」)。
補充為什麼用切換率,而不是直接算 KL
Markov toy 的狀態空間是 ,理論上可以查表算精確的 KL,但用有限樣本估 KL 會有很大的偏差(沒出現過的序列該算多少?),在 小的時候那個偏差會蓋掉真正的訊號。
切換率是這個資料分佈的一個充分好用的摘要:它有一個已知的真值(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,步數 各抽 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,步數 ,畫 kl_to_data。期待一條隨步數下降、漸變的曲線。把這張圖留好,U5 加上 remasking 之後會在同一張圖上疊新的線。
步驟 4:absorbing vs uniform 的修正實驗(圖 c)
另訓一個 uniform 鏈的模型(q_sample 改成「以 的機率換成均勻隨機字」,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 步、,所以「每步平均翻開一格」。但實際跑出來的合法率是 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 之下第 步的翻開機率是 ,而這一組機率等價於每個位置各自獨立抽一個 上的均勻標籤——位置 第一次被翻開在第 步的機率是 ,和 無關。於是「最後一次只翻一格」就是「最大的那個標籤只有一個位置拿到」:
代進去,、合法率 78.6%。所以量到 78–79% 不是實作沒調好——那就是這個設定的天花板,剩下的 21.4% 全部是因子化誤差。
三種修法,各對應不同的意思:
- 強制每步剛好一格(在還被遮的位置裡隨機挑一個)。這樣因子化誤差嚴格是 0,合法率會是 100%——這就是 U4.3 的 demo 做的事,也等於一個任意順序的 autoregressive 模型。
- 步數加大(例如 32 步)。翻開機率變小、最後一次只翻一格的機會上升,合法率往 100% 靠:同一條公式給出 是 88.4%、 是 94.0%、 是 96.9%。這是「多走幾步」的標準答案,而且收斂得比想像中慢。
- 什麼都不改,但把它記在報告裡。 這才是這一題真正的用意:「步數 」不等於「每步一格」,而 masked diffusion 的論文報數字時如果只寫步數、沒寫每步實際翻開幾格,那個數字就不能直接和 autoregressive 比。U4.3 引的 Zheng et al. 提醒的正是這件事。
作業
- Uniform vs absorbing 的合法率。在 parity toy 上對兩條鏈各畫步數 vs 合法率。uniform 鏈在 8 步時是否也接近 100%?若不是,說明是哪一種誤差(提示:uniform 鏈的 裡「只剩一個不確定的位置」這個狀態,模型認得出來嗎?)。
- 信心排序。把步驟 3 的
flip改成「每步翻開 marginal 最尖(entropy 最低)的 個位置」(MaskGIT 式),重畫 Markov toy 的曲線。改善多少?在 parity toy 上呢——為什麼那裡幾乎沒有幫助? - 換 parametrization。把網路改成直接預測 (而非 ),用 U4.1 的一般 ELBO 訓練。比較:(a) 訓練是否更難收斂;(b) 能不能跳步取樣。寫下你觀察到的差別與 U4.1 Q1 的對應。
- (選)拿掉 。把網路的 輸入移除、重訓 absorbing 模型。loss 曲線與步驟 3 的合法率是否改變?用 U4.2 那個
<Details>裡的論證解釋為什麼。
參考文獻
- Sahoo, S. S. et al. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.(步驟 1 的 loss 與步驟 3 的取樣程序。)
- Austin, J. et al. Structured Denoising Diffusion Models in Discrete State-Spaces. NeurIPS 2021.(步驟 4 uniform 鏈的一般 loss。)