U6.6 14 分鐘閱讀 2026年9月

U6.6 實作:三種「一步」放在同一張圖上

本篇重用M6.1兩個城市的人口分佈差多少:W₂ 距離·M0.3漂流的溫度計:全導數、微分穿過積分與 JVP

設定

沿用 toy.py(月牙資料、MLP、exact_score)與 U3.6fm.pycouplings.py(reflow×2 的權重直接載入)。本單元新增 cm.py:VE 座標下的加噪、時間網格、參數化、CD / CT 損失、一步與多步取樣。

三件事要先定下來。座標:本單元是 xt=x0+tϵx_t=x_0+t\epsilont[ε,T]t\in[\varepsilon,T]ε=0.002\varepsilon=0.002;資料已標準化,toy 上 T=10T=10 就夠遠(xTx_T 幾乎是純噪聲)。teacher:2D toy 的好處在這裡最明顯——真實 score 可以精確算,所以 CD 可以用 exact_score 當 teacher,把「teacher 誤差」這一項從實驗裡拿掉,剩下的就只有 Euler 的曲率誤差。作業會再用 U1.5 訓好的 ϵθ\epsilon_\theta 當 teacher 對照。評估:所有方法都用 2000 個樣本對資料的 W2W_2ot.emd2),與 U3.6 圖 b 同一把尺。

步驟 1:參數化與時間網格

EPS, T, SIGMA_D = 0.002, 10.0, 1.0          # 資料已標準化,σ_d = 1

def karras_grid(N, rho=7.0):                # t_1 = EPS < … < t_N = T,小 t 處密
    i = torch.arange(N)
    return (EPS**(1/rho) + i/(N-1) * (T**(1/rho) - EPS**(1/rho)))**rho

def c_skip(t): return SIGMA_D**2 / ((t - EPS)**2 + SIGMA_D**2)
def c_out(t):  return SIGMA_D * (t - EPS) / (SIGMA_D**2 + t**2).sqrt()

class ConsistencyModel(nn.Module):
    def __init__(self, F): super().__init__(); self.F = F      # F:toy.py 的 MLP
    def forward(self, x, t):
        return c_skip(t)[:, None] * x + c_out(t)[:, None] * self.F(x, t)

檢查點:cm(x, torch.full((n,), EPS)) 應與 x 逐元素相等到浮點誤差——邊界條件是架構給的,訓練前就成立(U6.2)。再算 c_skip(karras_grid(20)),確認它從 1 單調降到接近 0。

步驟 2:CD 與 CT 的訓練迴圈

def exact_score_ve(x, t):                    # VE 座標的精確 score:p_t = data ⊛ N(0, t²)
    ...                                      # 參考解答:改寫 toy.py 的混合 Gaussian 公式,σ_t² = t²

def train_cm(mode, N, steps, teacher=exact_score_ve):
    cm, tgt = ConsistencyModel(MLP()), None
    grid = karras_grid(N)
    for step in range(steps):
        x0  = sample_data(batch)
        n   = torch.randint(0, N - 1, (batch,))          # 相鄰的一對 (t_n, t_{n+1})
        tn, tn1 = grid[n], grid[n + 1]
        eps = torch.randn_like(x0)
        xt1 = x0 + tn1[:, None] * eps
        if mode == 'CD':                                  # teacher 走一步 Euler
            x_hat = xt1 + ((tn1 - tn) * tn1)[:, None] * teacher(xt1, tn1)
        else:                                             # CT:同一組 (x0, eps) 重加噪
            x_hat = x0 + tn[:, None] * eps
        with torch.no_grad():
            target = cm(x_hat, tn)                        # θ⁻ = stopgrad(θ),不用 EMA(iCT)
        loss = (huber(cm(xt1, tn1), target) / (tn1 - tn)).mean()   # λ = 1/Δt
        loss.backward(); opt.step(); opt.zero_grad()
    return cm

U1.5 的五行對照:抽 x0x_0、抽時間、加噪聲三行沒變;變的是目標——不再是 ϵ\epsilon,而是「同一條軌跡上前一格的點經過 fθf_{\theta^-} 的輸出」——以及 CD 與 CT 之間只差 x_hat 那一行。這一行就是 U6.3 的全部。

huber 是 pseudo-Huber xy2+c2c\sqrt{\|x-y\|^2+c^2}-cc=0.000542c=0.00054\sqrt{2};先用固定 N=40N=40 各訓 CD 與 CT 一個模型。檢查點:訓練中隨機抓一條精確軌跡(用 exact_score_ve 細步積分),沿軌跡取 8 個 (xt,t)(x_t,t)cm(x_t, t) 的輸出應幾乎相同,且接近軌跡終點——這就是 U6.2 核心 demo 裡那條「平線」。CD 的平線通常比 CT 更平、更早出現。

展開細節CT 的 N curriculum:如果固定 N 訓不起來

固定 N=40N=40 的 CT 在 toy 上通常可以直接收斂;若 loss 卡住或樣本糊成一團,把 N 改成隨訓練步數翻倍的排程(U6.4N(k)=min(s02k/K,s1)+1N(k)=\min(s_0 2^{\lfloor k/K'\rfloor},s_1)+1,toy 上 s0=4s_0=4s1=128s_1=128),並確認 target 那一行沒有用 EMA。這兩件事是 iCT 的第一與第五項。

步驟 3:CT 的 bias 對 NN(圖 a)

固定訓練步數與種子,對 N{2,4,8,16,32,64,128,256}N\in\{2,4,8,16,32,64,128,256\} 各訓一個 CT 模型與一個 CD 模型(teacher 是精確 score)。每個模型抽 2000 個 xTx_T、算一步 fθ(xT,T)f_\theta(x_T,T)W2W_2。兩條曲線畫在同一張圖上,橫軸 NN 對數刻度。

預期:CDNN 單調下降後趨平——趨平的位置是 MLP 的逼近極限;下降的部分是 Euler 一步的曲率誤差 Δt\propto\Delta tU6.3 的定理:fθf=O(Δt)\|f_\theta-f\|=O(\Delta t))。CT 先降後NN 小時 bias 主導(每一段被當成直線),NN 大時 variance 主導(相鄰兩點的訊號被條件目標的散佈淹沒),最低點在中間。兩條線在小 NN 幾乎重合——那裡兩者的主要誤差都是同一個 Euler 曲率項。

把 CT 曲線的最低點記下來,它是步驟 5 主圖裡 CT 的代表。

步驟 4:條件目標的扇子(圖 b)

這一步不訓練,直接驗證 U6.3 Q1 的等式。取一個 xtn+1x_{t_{n+1}},從後驗 p(x0xtn+1)p(x_0\mid x_{t_{n+1}})(混合 Gaussian,可精確抽)抽 500 個 x0x_0,每個算 ϵ=(xtn+1x0)/tn+1\epsilon=(x_{t_{n+1}}-x_0)/t_{n+1} 與 CT 目標 x0+tnϵx_0+t_n\epsilon。畫這 500 個點的散佈,疊上它們的質心、CD 的 Euler 點、以及真實的 xtnx_{t_n}(細步積分)。

xt1 = torch.tensor([[0.3, -0.2]]); tn1, tn = 2.0, 1.5
x0s   = posterior_sample(xt1, tn1, n=500)              # 混合 Gaussian 後驗
eps   = (xt1 - x0s) / tn1
ct    = x0s + tn * eps
euler = xt1 + (tn1 - tn) * tn1 * exact_score_ve(xt1, torch.tensor([tn1]))
print((ct.mean(0) - euler).norm())                      # 應 ≲ 5e-3(500 組的抽樣誤差)

對三個 Δt{1.0,0.3,0.1}\Delta t\in\{1.0,0.3,0.1\} 各畫一張:扇子的半徑 Δt\propto\Delta t(variance),Euler 點與真實點的距離 Δt2\propto\Delta t^2(bias)。把兩個量對 Δt\Delta t 畫成 log-log,斜率應分別接近 1 與 2。

步驟 5:三種「一步」同一張圖(圖 c,主圖)

三個一步生成器、同一組 2000 個起點噪聲:

  • CDN=128N=128、精確 score teacher):fθ(xT,T)f_\theta(x_T,T)
  • CT(步驟 3 的最佳 NN):fθ(xT,T)f_\theta(x_T,T)
  • reflow×2U3.6 載入):一步 Euler。注意座標轉換——reflow 模型活在 FM 座標 s[0,1]s\in[0,1],起點是 N(0,I)\mathcal N(0,I);用 xT/(1+T)x_T/(1+T) 當它的起點(U6.0 的對照表:xsFM=sxtx^{\text{FM}}_s=s\,x_tt=1sst=\frac{1-s}{s},所以 t=Tt=T 對應 s=11+Ts=\frac1{1+T}xsFM=xT/(1+T)x^{\text{FM}}_s=x_T/(1+T)T=10T=10 時寫成 xT/Tx_T/T 會差 10%)。

畫三組樣本並排(同一座標範圍)、下方一張 W2W_2 長條圖。預期:CD 最好(teacher 是精確的、NN 大);CT 略差(variance 極限);reflow×2 與 CT 接近或稍差(殘餘曲率)。三根長條旁各寫一句「誤差來源」:CD——MLP 逼近極限;CT——bias–variance 極限;reflow×2——x¨dt\int\|\ddot x\|dt 尚未為零。

把這張圖留好,U7 會加上 MeanFlow 的一步、再加上多步版本。

課堂提問Q1

步驟 5 把 CD、CT、reflow×2 放在同一張長條圖上比。但這三個東西拿到的資源完全不同:CD 有一個精確的 score teacher,CT 完全不用 teacher,reflow×2 用的是一個訓出來的速度場、而且已經跑過兩輪 reflow。

這樣比,比得出什麼?比不出什麼?

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

比得出的是**「在這個 toy 上,拿這些資源能做到多好」**——一個工程上的對照。比不出的是「哪一個方法比較好」,因為三者的輸入根本不是同一組。

U2.4 那四個旋鈕的教訓在這裡完全適用:一個實驗結果要歸因到哪一個因素,得把其他因素控制住才說得清。 這張圖至少有三個沒控制住的東西:

  • teacher 的品質。 CD 的上限就是 teacher;給它精確 score,等於給它一個真實模型上拿不到的上限。作業 1 把 teacher 換成訓好的 ϵθ\epsilon_\theta 再重畫,就是在控制這一項——CD 的長條會升高多少,那個差就是「teacher 誤差」。
  • 訓練預算。 reflow×2 背後是「訓一個 FM + 生一次資料 + 再訓一次 + 再生一次 + 再訓一次」,總計算量比 CD 或 CT 的一次訓練多得多。三根長條下面應該另外標各自的總網路呼叫次數。
  • 配對。 reflow×2 的 (x0,ϵ)(x_0,\epsilon) 是拉直過的耦合,CD 與 CT 用的是獨立配對。作業 3 正是把這一項搬過去。

那這張圖還有什麼用?兩件事。第一,它讓「一步生成」這件事在同一把尺上有了數字——在此之前 U6.0U6.4 講的都是機制,沒有量。第二,它把每一根長條的誤差來源標出來了(MLP 逼近極限/bias–variance 極限/殘餘曲率),而那三個來源彼此不同,這本身就是結論:一步生成沒有單一瓶頸。

寫報告時把上面三項「沒控制住」的東西明白寫出來,比多跑一組實驗更有價值。

步驟 6:多步 CM 還吃不吃步數?(圖 d)

@torch.no_grad()
def multistep_cm(cm, xT, taus):                 # taus:遞減的中間時刻,例如 [T, 3.0, 1.0, 0.3]
    x0_hat = cm(xT, torch.full((len(xT),), taus[0]))
    for tau in taus[1:]:
        x_tau  = x0_hat + tau * torch.randn_like(x0_hat)     # 重加噪:新的噪聲
        x0_hat = cm(x_tau, torch.full((len(xT),), tau))
    return x0_hat

對 CT 模型,步數 k{1,2,4,8,16,32,64}k\in\{1,2,4,8,16,32,64\}τ\tau 序列用 karras_grid(k+1) 反轉),畫距離對 kk;同一張圖疊上 U3.6 那個獨立配對 FM 的 Euler kk 步。FM 那條線是斜率約 1-1 的直線(U3.2O(1/N)O(1/N))。CM 那條線跟不跟得上那個斜率,就是這一步要看的事——U6.5 Q1 的三個面說它不該跟得上,但「不跟得上」有好幾種長法,量到哪一種就寫哪一種,kk 要掃到 64 才看得出來。

再做一件事:把 cm 換成真值 ff(用 exact_score_ve 細步積分到 ε\varepsilon),同樣的多步程序在所有 kk 都應給接近零的距離——這就是 U6.5w8-5-c,也確認了差別全部來自 fθf_\theta

距離不要用 W2W_2 這張圖上 CM 與真值 ff 的差距會小到 10310^{-3} 量級,而 n=1000n=1000 的經驗 W2W_2 在這個 toy 上光是「兩組真資料互比」就有 0.3 的地板,訊號整個被蓋掉。改用無偏的 energy distanceiji\ne j 的 U-statistic,U3.6 用過),n=8000n=8000,地板在 10410^{-4}。同一顆種子跑三次取平均、把全距也畫上去。

作業

  1. Teacher 換成網路。 把步驟 2 的 CD teacher 從 exact_score_ve 換成 U1.5 訓好的 ϵθ\epsilon_\theta(轉到 VE 座標:logptϵθ/t\nabla\log p_t\approx-\epsilon_\theta/t,注意 VP→VE 的時間對應)。重畫圖 a 的 CD 曲線。它在大 NN 處的平台是升高還是不變?說明哪一項誤差進來了。
  2. EMA target。 在步驟 2 加回 θ=EMA0.999(θ)\theta^-=\text{EMA}_{0.999}(\theta),只對 CT 重畫圖 a。曲線的最低點怎麼移動?對照 U6.4 對這一項的分類(bias)。
  3. 先拉直、再 CT。 用 reflow×2 的耦合(不是獨立配對)造 (x0,ϵ)(x_0,\epsilon)——也就是把 U3.6(x0, x1) 配對換成本單元座標的 (x0,ϵ)(x_0,\epsilon)——重訓 CT。圖 a 的 CT 曲線在小 NN 處的 bias 下降了多少?用 U6.3 那個「CT 的目標是假設這一段是直線」的展開框解釋。
  4. 連續時間(sCT)。torch.func.jvp 實作 U6.4 的全導數目標(toy 上直接用 VE 座標即可,不必換 TrigFlow):條件速度用 ϵ\epsilon,tangent 除以 +0.1\|\cdot\|+0.1。與離散 CT 的最佳 NN 比一步 W2W_2。再把 tangent normalization 關掉重訓一次,觀察 loss 曲線。
  5. (選)最佳 τ1\tau_1 對兩步 CM 掃 τ1[0.1,T]\tau_1\in[0.1,T],畫 W2W_2τ1\tau_1。最低點在哪裡?τ1ε\tau_1\to\varepsilonτ1T\tau_1\to T 兩端各退化成什麼?
想一想

步驟 3 裡 CT 的 W2W_2 曲線在大 NN上升,而 CD 的不會。原因是:

想一想

步驟 6 中把 fθf_\theta 換成真值 ff,多步 CM 在所有步數都給接近零的距離。這說明多步 CM 那條線的形狀來自:

參考文獻

  1. Song, Y., Dhariwal, P., Chen, M., Sutskever, I. Consistency Models. ICML 2023.(步驟 2 的 CD / CT 損失、步驟 6 的多步取樣。)
  2. Song, Y., Dhariwal, P. Improved Techniques for Training Consistency Models. ICLR 2024.(pseudo-Huber、λ=1/Δt\lambda=1/\Delta t、無 EMA target、NN curriculum。)
  3. Karras, T., Aittala, M., Aila, T., Laine, S. Elucidating the Design Space of Diffusion-Based Generative Models. NeurIPS 2022.(karras_gridcskip,coutc_{\text{skip}},c_{\text{out}}。)
  4. Liu, X., Gong, C., Liu, Q. Flow Straight and Fast. ICLR 2023.(步驟 5 的 reflow×2 與作業 3 的耦合。)