U5.6 15 分鐘閱讀 2026年9月

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.5toy_discrete.py 上:parity toy(L=8L=8)、Markov toy(L=16L=16)、exact_conditional(xt)legal(x)kl_to_data(samples)。這個單元多兩個精確工具,都用查表做:

  • exact_marginal(xt, t)pt(xt)=x0qt(xtx0)p(x0)p_t(x_t)=\sum_{x_0}q_t(x_t\mid x_0)\,p(x_0)
  • exact_ratios(xt, t):對 xtx_t 的每個鄰居 yypt(y)/pt(xt)p_t(y)/p_t(x_t)——U5.2 說它等於 exact_conditionalαt/(1αt)\alpha_t/(1-\alpha_t)(absorbing),這是第一個要用程式驗證的等式。

時間改成連續 t[0,1]t\in[0,1],線性 schedule αt=1t\alpha_t=1-t,forward rate βt=1/(1t)\beta_t=1/(1-t),翻開 rate ut=α˙t/(1αt)=1/tu_t=-\dot\alpha_t/(1-\alpha_t)=1/t。上一個單元的網路吃整數 t{1,,32}t\in\{1,\dots,32\},這個單元要用它時餵 round(t*32);或者用下面步驟 3 的方式重訓一個吃連續 tt 的版本(loss 一模一樣,只是 tUniform[0,1]t\sim\mathrm{Uniform}[0,1]、權重 1/t1/t)。

步驟 1:τ-leaping 取樣器(圖 a)

CTMC 的精確模擬是 Gillespie 演算法:抽下一次跳躍的時間與位置,一次只跳一個 token。這太慢(LL 個位置、幾千次跳躍)。τ-leaping [1, 2] 是 CTMC 的 Euler 法:固定一段 τ\tau,用步首的 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 逐行對照:那裡翻開的機率是 αˉsαˉt1αˉt\frac{\bar\alpha_s-\bar\alpha_t}{1-\bar\alpha_t},這裡是 1eutτ1-e^{-u_t\tau}。兩者在 τ0\tau\to0 時一致;在有限 τ\tau 下前者是精確的跳步(用了 Qˉ\bar Q 的 closed form),後者是 Euler 近似。圖 a:對 parity toy 畫步數 {1,2,4,8,16,32}\{1,2,4,8,16,32\} 的合法率,兩個取樣器各一條線。期待:兩條線很接近但不重合,而且都離 100% 還有一段距離——4/8/32 步大致落在 57–65 / 73–79 / 93–95%(下面的 demo 量的就是這條線),1、2 步更低。精確跳步那條線有標準答案可以對:U4.5 Q1 那條 (1+q)/2(1+q)/2 在 4/8/32 步分別給出 64.1 / 78.6 / 94.0%。差距集中在步數最少的那幾個點。

檢查點:τ-leaping 裡「所有位置獨立決定」那一行,就是因子化誤差的來源;把 steps 開到 256(每步幾乎最多跳一個位置),兩條線都會到 100%。同一個 toy、同一個網路,換了語言,得到差不多同一張圖——這是這個單元想讓你確認的第一件事,而「差多少」正好是下面 demo 要量的。

補充兩端的 rate 會發散,程式裡要怎麼收

線性 schedule 下翻開 rate 是 ut=1/tu_t=1/tt0t\to0 時發散;forward rate βt=1/(1t)\beta_t=1/(1-t)t1t\to1 發散。這不是 bug——U5.1 說過它的意思是「剩下的少數 [MASK] 要以很快的速度翻開」。

程式裡有兩種收法:把最後一步的 tt clamp 在一個小的 tmint_{\min}(例如 10310^{-3}),或者直接規定最後一步把所有殘留的 [MASK]pθp_\theta 一次填完。兩種都會引入一點偏差;後者比較常見,因為它保證輸出裡沒有 [MASK]這一項偏差和步數無關,所以在畫「步數 vs 合法率」時它是一條水平的地板,不要把它算成因子化誤差。

互動 demo:兩個取樣器差多少。 parity toy(L=8L=8),橫軸步數、縱軸合法率,每個點 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 個百分點,而且差距集中在步數最少的那幾個點——那裡的 τ\tau 最大、一階近似最粗。切到「強制每步剛好一格」差別大得多:精確跳步在 8 步(=L=L)就衝到 100%,τ-leaping 只有 71–72%,因為它的每步機率不同,有些步一格都不開、最後被迫一次補完好幾格。每個點 2000 條樣本,數字每次跑都會動一兩個百分點。

課堂提問Q1

上面量到精確跳步在每一個步數上都不輸 τ-leaping,而且它不需要 rate——只要 αˉ\bar\alpha 的 closed form 就能算每步翻開的機率。那我們為什麼還要學 τ-leaping?

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

會被想到的回答大概是「τ-leaping 比較快」或「精確跳步只是 τ-leaping 的特例」。兩個都不對:兩者每步都只呼叫網路一次,成本一樣;而精確跳步用的是closed form,不是把 τ\tau 取小的極限。

真正的理由是closed form 只對「寫得出 Qˉt\bar Q_t」的鏈存在。這一整個單元在做的事都在破壞那個前提:

  • 加上 remasking 之後,反向鏈的 rate 依賴 σt\sigma_t 與網路輸出,沒有 Qˉ\bar Q 可以查(U5.3)。
  • 逐位置不同的 source 與 κt\kappa_t^\ell、混合路徑(U5.4),一樣沒有一個統一的 closed form。
  • 用比值直接寫反向 rate 時(步驟 3),我們手上只有 rate,本來就沒有機率的 closed form。

τ-leaping 只需要「現在每根管子的 rate 是多少」,所以上面每一種情形它都能跑。rate 是比 Qˉ\bar Q 更通用的介面,這就是換語言換到的東西。

上面的數字還有一件事要讀出來:精確跳步的「精確」是逐位置的精確——每個位置在 [s,t][s,t] 之間翻開的機率算得完全對。但它在 4 步時也只有 63–65%,離 100% 很遠。缺的那一塊不是機率算錯,是同時翻開的位置被獨立抽U4.3)。所以「精確跳步」精確的是邊際,不是 joint;兩個取樣器的差距(5–8 個百分點)比因子化誤差(4 步的 63% 到 32 步的 95%,三十幾個百分點)小得多,這個數量級的對比本身就是步驟 1 的結論。

步驟 2:remasking 強度 vs 合法率(圖 b,疊在上一個單元的圖上)

上面的 sigma 參數已經把 U5.3 寫進去了:已翻開位置以 rate σ\sigma 遮回去,被遮位置的翻開 rate 加大 σαt/(1αt)\sigma\alpha_t/(1-\alpha_t)。先做邊際檢查:對 Markov toy 跑 2000 條軌跡到 t=0.5t=0.5,數 [MASK] 的比例,對 σ{0,0.5,1,2}\sigma\in\{0,0.5,1,2\} 應該都是 0.5±0.5\pm 抽樣誤差;把補償那一行拿掉(u_total = u)再跑一次,比例會隨 σ\sigma 明顯上升。這兩個數字寫進 notebook。

然後畫圖 b:

  • 對 Markov toy,橫軸步數 {4,8,16,32,64,128}\{4,8,16,32,64,128\}(對數刻度,與上一個單元相同),縱軸 kl_to_data,對 σ{0,0.5,1,2}\sigma\in\{0,0.5,1,2\} 各一條線。σ=0\sigma=0 那條就是上一個單元的圖 b,把它疊回去。
  • 對 parity toy 畫同樣的圖,縱軸換成合法率。

先寫下你的猜測再跑σ>0\sigma>0 的線會在哪些步數壓到 σ=0\sigma=0 底下?把預測寫在 notebook 裡再看結果。

我們在 U5.3 的 demo 上量到的是:用 exact_conditional 當網路時,只有步數少的那一端有明顯好處(Markov toy 的切換率在 N=8N=8 從 0.14–0.15 掉到 0.13 左右),步數多的一端 σ=0\sigma=0 本來就已經打在資料的真值上,沒有空間可以進步。理由是網路沒有誤差時唯一的誤差來源就是因子化,而它隨步數自己變小。

所以這一步兩種網路都要跑exact_conditional 與訓出來的網路,把兩組「σ\sigma vs 指標」畫在一起,看有多少改善來自「修掉因子化的錯」、多少來自「修掉網路的錯」。ReMDM [4] 說的 inference-time scaling 是在真實模型上量的,那裡剩下的誤差是網路在 context 不完整時猜錯——要在 toy 上看到它,得先讓網路真的會錯。

步驟 3:SEDD 的比值學習 vs MDLM 的 cross-entropy(圖 c)

網路換成輸出比值:對每個位置輸出 KK 個正數(用 expsoftplus),代表 sθ(xt),vpt(xt(v))/pt(xt)s_\theta(x_t)_{\ell,v}\approx p_t(x_t^{(\ell\to v)})/p_t(x_t)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()

訓好後在一批 (xt,t)(x_t,t) 上做三件事:

  1. 比值 vs 精確比值sexact_ratios 在被遮位置的相對誤差,對 tt 畫曲線。
  2. 比值 vs MDLM 網路:把上一個單元(或步驟 1 重訓)的 model(xt,t).softmax(-1)αt/(1αt)\alpha_t/(1-\alpha_t),與 s 比對。期待:兩者在資料密集區幾乎重合——同一組數字、兩種 parametrization
  3. 取樣:用比值直接寫反向 rate Rˉt([MASK]v)=βtsθ(xt),v\bar R_t(\texttt{[MASK]}\to v)=\beta_t\,s_\theta(x_t)_{\ell,v} 做 τ-leaping,與步驟 1 的曲線比。

圖 c:三個子圖並排。若 (2) 的差距明顯大於 (1) 的誤差,檢查 aa 的常數與權重——最常見的 bug 是把 αt/(1αt)\alpha_t/(1-\alpha_t) 寫反。

作業

  1. Uniform 鏈上的比值。 對 uniform 鏈推出條件比值 qt(yx0)/qt(xx0)q_t(y\mid x_0)/q_t(x\mid x_0)(位置 \ell 從字 xx^\ell 換成 yy^\ell;提示:只剩位置 \ell 的因子,分子分母各是 αt1[=x0]+(1αt)/K\alpha_t\mathbb 1[\cdot=x_0^\ell]+(1-\alpha_t)/K),寫出它的後驗平均是 p(x0xt)p(x_0^\ell\mid x_t) 的什麼函數。訓一個 uniform 鏈的比值網路,與 U4.5 作業 1 的 D3PM 網路比較合法率。哪一個 target 在 uniform 鏈上訓得比較穩?
  2. σt\sigma_t 的 schedule。 把常數 σ\sigma 換成鐘形(中段大、兩端小),在圖 b 上再疊一條線。與常數 σ\sigma 在相同總「重遮量」σt(1mt)dt\int\sigma_t(1-m_t)\,dtmtm_t[MASK] 的比例,見 U5.3)下比較。
  3. 信心排序 vs remasking。U4.5 作業 2 的信心排序與這個單元的 σ>0\sigma>0 放在同一張圖。信心排序不保持邊際、remasking 保持——在 toy 上哪一個更好?用步驟 2 的邊際檢查量化前者偏了多少。
  4. (選)精確模擬。 實作 Gillespie 演算法(一次跳一個 token),確認它在 parity toy 上的合法率是 100%、且與 τ-leaping 在 steps 很大時一致。它就是 U4.3 說的「任意順序 AR」。
想一想

τ-leaping 一步裡「所有位置獨立決定要不要跳、跳到哪」這一行,對應的是:

想一想

步驟 3 裡 SEDD 網路輸出的比值與 MDLM 網路輸出的 pθ(x0xt)p_\theta(x_0^\ell\mid x_t) 在 absorbing 鏈上「幾乎重合」,重合的方式是:

參考文獻

  1. Campbell, A., Benton, J., De Bortoli, V., Rainforth, T., Deligiannidis, G., Doucet, A. A Continuous Time Framework for Discrete Denoising Models. NeurIPS 2022.(τ-leaping 用於離散擴散取樣。)
  2. Gillespie, D. T. Approximate Accelerated Stochastic Simulation of Chemically Reacting Systems. J. Chem. Phys. 115, 1716 (2001).(τ-leaping 的出處。)
  3. Lou, A., Meng, C., Ermon, S. Discrete Diffusion Modeling by Estimating the Ratios of the Data Distribution. ICML 2024.(步驟 3 的 loss。)
  4. Wang, G., Schiff, Y., Sahoo, S. S., Kuleshov, V. Remasking Discrete Diffusion Models with Inference-Time Scaling. 2025.(步驟 2。)
  5. Sahoo, S. S. et al. Simple and Effective Masked Diffusion Language Models. NeurIPS 2024.(步驟 1 的 MDLM 網路。)