U5.6 實作:τ-leaping、Remasking 與比值學習
本篇重用M4.2電話隨時會響:Continuous-Time Markov Chain 與 Rate Matrix·M4.3影片倒著播看得出來嗎:Time Reversal 與比值·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy
上一個單元的取樣器和這個單元的 rate 語言,
跑出來是同一條曲線嗎?
沿用上一個單元的 toy
一切都建在 U4.5 的 toy_discrete.py 上:parity toy()、Markov toy()、exact_conditional(xt)、legal(x)、kl_to_data(samples)。這個單元多兩個精確工具,都用查表做:
exact_marginal(xt, t):。exact_ratios(xt, t):對 的每個鄰居 算 ——U5.2 說它等於exact_conditional乘 (absorbing),這是第一個要用程式驗證的等式。
時間改成連續 ,線性 schedule ,forward rate ,翻開 rate 。上一個單元的網路吃整數 ,這個單元要用它時餵 round(t*32);或者用下面步驟 3 的方式重訓一個吃連續 的版本(loss 一模一樣,只是 、權重 )。
步驟 1:τ-leaping 取樣器(圖 a)
CTMC 的精確模擬是 Gillespie 演算法:抽下一次跳躍的時間與位置,一次只跳一個 token。這太慢( 個位置、幾千次跳躍)。τ-leaping [1, 2] 是 CTMC 的 Euler 法:固定一段 ,用步首的 rate 算每個位置在這段時間內跳的機率,所有位置獨立地決定要不要跳、跳到哪。
@torch.no_grad()
def tau_leap_step(model, xt, t, tau, sigma=0.0):
n, L = xt.shape
probs = model(xt, t).softmax(-1) # p_θ(x0^ℓ | x_t), (n, L, K)
alpha = 1 - t
u = 1 / t # 翻開 rate −α̇/(1−α)
u_total = u + sigma * alpha / (1 - alpha) # remasking 的補償(步驟 2)
masked = xt == MASK
# 被遮位置:以機率 1−exp(−u_total τ) 翻開,翻成哪個字按 p_θ 抽
flip = masked & (torch.rand(n, L) < 1 - torch.exp(-u_total * tau))
new = torch.distributions.Categorical(probs).sample()
xt = torch.where(flip, new, xt)
# 已翻開位置:以機率 1−exp(−σ τ) 遮回去(步驟 2;σ=0 時不動)
remask = (~masked) & (torch.rand(n, L) < 1 - torch.exp(-sigma * tau))
return torch.where(remask, torch.full_like(xt, MASK), xt)
def sample_ctmc(model, n, steps, sigma=0.0):
xt = torch.full((n, L), MASK)
ts = torch.linspace(1, 0, steps + 1) # t: 1 → 0
for t, s in zip(ts[:-1], ts[1:]):
xt = tau_leap_step(model, xt, t, tau=t - s, sigma=sigma)
return xt # 最後一步 t→0 時 u→∞,記得把殘留的 [MASK] 全部翻開
與 U4.5 步驟 3 的 sample 逐行對照:那裡翻開的機率是 ,這裡是 。兩者在 時一致;在有限 下前者是精確的跳步(用了 的 closed form),後者是 Euler 近似。圖 a:對 parity toy 畫步數 的合法率,兩個取樣器各一條線。期待:兩條線很接近但不重合,而且都離 100% 還有一段距離——4/8/32 步大致落在 57–65 / 73–79 / 93–95%(下面的 demo 量的就是這條線),1、2 步更低。精確跳步那條線有標準答案可以對:U4.5 Q1 那條 在 4/8/32 步分別給出 64.1 / 78.6 / 94.0%。差距集中在步數最少的那幾個點。
檢查點:τ-leaping 裡「所有位置獨立決定」那一行,就是因子化誤差的來源;把 steps 開到 256(每步幾乎最多跳一個位置),兩條線都會到 100%。同一個 toy、同一個網路,換了語言,得到差不多同一張圖——這是這個單元想讓你確認的第一件事,而「差多少」正好是下面 demo 要量的。
補充兩端的 rate 會發散,程式裡要怎麼收
線性 schedule 下翻開 rate 是 , 時發散;forward rate 在 發散。這不是 bug——U5.1 說過它的意思是「剩下的少數 [MASK] 要以很快的速度翻開」。
程式裡有兩種收法:把最後一步的 clamp 在一個小的 (例如 ),或者直接規定最後一步把所有殘留的 [MASK] 依 一次填完。兩種都會引入一點偏差;後者比較常見,因為它保證輸出裡沒有 [MASK]。這一項偏差和步數無關,所以在畫「步數 vs 合法率」時它是一條水平的地板,不要把它算成因子化誤差。
互動 demo:兩個取樣器差多少。 parity toy(),橫軸步數、縱軸合法率,每個點 2000 條樣本。兩條線很接近但不重合:4/8/32 步的合法率是精確跳步 63–65 / 78–79 / 94–95%(理論值 64.1 / 78.6 / 94.0%,三個點全部對上)、τ-leaping 57–59 / 73–74 / 93–94%,最大差 5–8 個百分點,而且差距集中在步數最少的那幾個點——那裡的 最大、一階近似最粗。切到「強制每步剛好一格」差別大得多:精確跳步在 8 步()就衝到 100%,τ-leaping 只有 71–72%,因為它的每步機率不同,有些步一格都不開、最後被迫一次補完好幾格。每個點 2000 條樣本,數字每次跑都會動一兩個百分點。
課堂提問Q1
上面量到精確跳步在每一個步數上都不輸 τ-leaping,而且它不需要 rate——只要 的 closed form 就能算每步翻開的機率。那我們為什麼還要學 τ-leaping?
先想一想,再展開看整理後的答案
會被想到的回答大概是「τ-leaping 比較快」或「精確跳步只是 τ-leaping 的特例」。兩個都不對:兩者每步都只呼叫網路一次,成本一樣;而精確跳步用的是closed form,不是把 取小的極限。
真正的理由是closed form 只對「寫得出 」的鏈存在。這一整個單元在做的事都在破壞那個前提:
- 加上 remasking 之後,反向鏈的 rate 依賴 與網路輸出,沒有 可以查(U5.3)。
- 逐位置不同的 source 與 、混合路徑(U5.4),一樣沒有一個統一的 closed form。
- 用比值直接寫反向 rate 時(步驟 3),我們手上只有 rate,本來就沒有機率的 closed form。
τ-leaping 只需要「現在每根管子的 rate 是多少」,所以上面每一種情形它都能跑。rate 是比 更通用的介面,這就是換語言換到的東西。
上面的數字還有一件事要讀出來:精確跳步的「精確」是逐位置的精確——每個位置在 之間翻開的機率算得完全對。但它在 4 步時也只有 63–65%,離 100% 很遠。缺的那一塊不是機率算錯,是同時翻開的位置被獨立抽(U4.3)。所以「精確跳步」精確的是邊際,不是 joint;兩個取樣器的差距(5–8 個百分點)比因子化誤差(4 步的 63% 到 32 步的 95%,三十幾個百分點)小得多,這個數量級的對比本身就是步驟 1 的結論。
步驟 2:remasking 強度 vs 合法率(圖 b,疊在上一個單元的圖上)
上面的 sigma 參數已經把 U5.3 寫進去了:已翻開位置以 rate 遮回去,被遮位置的翻開 rate 加大 。先做邊際檢查:對 Markov toy 跑 2000 條軌跡到 ,數 [MASK] 的比例,對 應該都是 抽樣誤差;把補償那一行拿掉(u_total = u)再跑一次,比例會隨 明顯上升。這兩個數字寫進 notebook。
然後畫圖 b:
- 對 Markov toy,橫軸步數 (對數刻度,與上一個單元相同),縱軸
kl_to_data,對 各一條線。 那條就是上一個單元的圖 b,把它疊回去。 - 對 parity toy 畫同樣的圖,縱軸換成合法率。
先寫下你的猜測再跑: 的線會在哪些步數壓到 底下?把預測寫在 notebook 裡再看結果。
我們在 U5.3 的 demo 上量到的是:用 exact_conditional 當網路時,只有步數少的那一端有明顯好處(Markov toy 的切換率在 從 0.14–0.15 掉到 0.13 左右),步數多的一端 本來就已經打在資料的真值上,沒有空間可以進步。理由是網路沒有誤差時唯一的誤差來源就是因子化,而它隨步數自己變小。
所以這一步兩種網路都要跑:exact_conditional 與訓出來的網路,把兩組「 vs 指標」畫在一起,看有多少改善來自「修掉因子化的錯」、多少來自「修掉網路的錯」。ReMDM [4] 說的 inference-time scaling 是在真實模型上量的,那裡剩下的誤差是網路在 context 不完整時猜錯——要在 toy 上看到它,得先讓網路真的會錯。
步驟 3:SEDD 的比值學習 vs MDLM 的 cross-entropy(圖 c)
網路換成輸出比值:對每個位置輸出 個正數(用 exp 或 softplus),代表 。U5.2 的 denoising score entropy,權重取 forward rate:
def score_entropy_loss(s, xt, x0, t): # s: (n, L, K) 正數
alpha = (1 - t)[:, None, None]
masked = (xt == MASK)[:, :, None] # absorbing 下只有被遮位置有鄰居
a = alpha / (1 - alpha) * F.one_hot(x0, K) # 條件比值 q(y|x0)/q(x|x0)
K_a = torch.where(a > 0, a * (a.log() - 1), torch.zeros_like(a))
bregman = s - a * s.log() + K_a # s − a log s + a(log a − 1)
w = (1 / (1 - t))[:, None, None] # 權重 = forward rate β_t
return (w * bregman * masked).sum((1, 2)).mean()
訓好後在一批 上做三件事:
- 比值 vs 精確比值:
s與exact_ratios在被遮位置的相對誤差,對 畫曲線。 - 比值 vs MDLM 網路:把上一個單元(或步驟 1 重訓)的
model(xt,t).softmax(-1)乘 ,與s比對。期待:兩者在資料密集區幾乎重合——同一組數字、兩種 parametrization。 - 取樣:用比值直接寫反向 rate 做 τ-leaping,與步驟 1 的曲線比。
圖 c:三個子圖並排。若 (2) 的差距明顯大於 (1) 的誤差,檢查 的常數與權重——最常見的 bug 是把 寫反。
作業
- Uniform 鏈上的比值。 對 uniform 鏈推出條件比值 (位置 從字 換成 ;提示:只剩位置 的因子,分子分母各是 ),寫出它的後驗平均是 的什麼函數。訓一個 uniform 鏈的比值網路,與 U4.5 作業 1 的 D3PM 網路比較合法率。哪一個 target 在 uniform 鏈上訓得比較穩?
- 的 schedule。 把常數 換成鐘形(中段大、兩端小),在圖 b 上再疊一條線。與常數 在相同總「重遮量」( 是
[MASK]的比例,見 U5.3)下比較。 - 信心排序 vs remasking。 把 U4.5 作業 2 的信心排序與這個單元的 放在同一張圖。信心排序不保持邊際、remasking 保持——在 toy 上哪一個更好?用步驟 2 的邊際檢查量化前者偏了多少。
- (選)精確模擬。 實作 Gillespie 演算法(一次跳一個 token),確認它在 parity toy 上的合法率是 100%、且與 τ-leaping 在
steps很大時一致。它就是 U4.3 說的「任意順序 AR」。
參考文獻
- Campbell, A., Benton, J., De Bortoli, V., Rainforth, T., Deligiannidis, G., Doucet, A. A Continuous Time Framework for Discrete Denoising Models. NeurIPS 2022.(τ-leaping 用於離散擴散取樣。)
- Gillespie, D. T. Approximate Accelerated Stochastic Simulation of Chemically Reacting Systems. J. Chem. Phys. 115, 1716 (2001).(τ-leaping 的出處。)
- Lou, A., Meng, C., Ermon, S. Discrete Diffusion Modeling by Estimating the Ratios of the Data Distribution. ICML 2024.(步驟 3 的 loss。)
- Wang, G., Schiff, Y., Sahoo, S. S., Kuleshov, V. Remasking Discrete Diffusion Models with Inference-Time Scaling. 2025.(步驟 2。)
- Sahoo, S. S. et al. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.(步驟 1 的 MDLM 網路。)