U3.6 13 分鐘閱讀 2026年9月

U3.6 實作:把曲率積分畫成一條線

本篇重用M2.3導航每 30 秒更新一次:Euler 法與它的誤差·M6.1兩個城市的人口分佈差多少:W₂ 距離

設定

沿用 toy.pyfm.py。這個單元新增三個檔:couplings.py(reflow 的配對生成、minibatch OT)、interpolant.py(帶 γt\gamma_t 的訓練與一整族 SDE 取樣器)、curvature.py。所有模型同架構、同 seed、同步數——只有配對不一樣。

步驟 1:四種配對

# (a) 獨立:flow matching 實作篇的模型,直接載入
# (b) reflow×1
x0 = torch.randn(N, 2)
x1 = euler_sample(model_indep, x0, n_steps=200)[-1]   # 用模型自己的 ODE 產生配對
model_reflow1 = train_fm(pairs=(x0, x1))               # 用這組固定配對重訓
# (c) reflow×2:對 model_reflow1 再做一次
# (d) minibatch OT
from scipy.optimize import linear_sum_assignment
def ot_pair(x0, x1):
    C = torch.cdist(x0, x1) ** 2
    _, col = linear_sum_assignment(C.numpy())
    return x0, x1[col]
model_ot = train_fm(pair_fn=ot_pair, batch=256)

訓練迴圈唯一的改動是 x0, x1 從哪裡來。其他一行都沒動。

步驟 2:軌跡圖(圖 a)

四個模型、相同的 32 個起點、Euler 10 步,並排畫。

會看到的事情是:由左到右越來越直,reflow×2 與 OT 已經接近直線,而殘餘的彎曲集中在兩個月牙之間——那裡是條件直線交叉最密的地方。

互動 demo:四種配對放在一起量。 這是圖 a 加上步驟 4 的表格的「標準答案版」——速度場精確算出來,所以沒有網路誤差混進去。四個指標的實測值:獨立配對 S ≈ 1.6、1 步誤差 ≈ 1.5、交叉比例 ≈ 23%;reflow×1 把 S 壓到 0.03、誤差 0.12、交叉 1%;batch OT(32)走到 S 0.22、誤差 0.30;全域 OT 是 S 0.02、誤差 0.06。reflow×1 和全域 OT 在「直」上打成平手,但運輸成本 OT 一貫略低——「直」和「省」真的是兩件事。

四種配對的曲率積分差好幾倍,而它們的兩端分布完全一樣——旋鈕動的是路,不是終點。 至於「曲率積分小,誤差是不是真的跟著小」,那要等下一步把兩個量畫在同一張圖上才知道。

步驟 3:曲率積分 vs 誤差(圖 b)

def curvature_integral(model, x0, n_fine=1000):
    traj = euler_sample(model, x0, n_fine)          # 細步當成真實軌跡
    v = (traj[1:] - traj[:-1]) * n_fine             # 速度 ≈ 差分
    a = (v[1:] - v[:-1]) * n_fine                   # 加速度
    return a.norm(dim=-1).mean(dim=0).mean()        # 對 t 積分、再對起點平均

對四個模型算 Ex¨dt\mathbb E\int\|\ddot x\|\,dt,以及 Euler 10 步終點的 W2W_2,畫成散點、兩軸取對數。

會看到的事情是:四個點近似共線、斜率約 1。這張圖就是 U3.2 那條不等式的實證——高度由曲率決定U3.2 的 demo 已經先用精確場把這條線畫出來過(六個設定、rr 通常 0.95 以上),這裡要看的是換成訓練出來的網路之後還成不成立

再加一組步數(N{5,10,20}N\in\lbrace 5,10,20 \rbrace):點會整體上下平移,但相對順序不變。

曲率積分小,誤差就一定小嗎?
這件事能不能在自己的圖上驗一次?

步驟 4:straightness 與運輸成本

對每個模型算 SSU3.3)與 EX1X02\mathbb E\|X_1-X_0\|^2,做成一張表。

會看到的事情是:reflow 的兩個數都單調下降,而 OT 的運輸成本最低。把 OT 的成本和全域 OT(2000 點精確解)比一下,就知道 reflow 的極限離 OT 有多遠。

步驟 5:ODE 與 SDE 在第幾步交叉(圖 c)

interpolant.py 訓一個 γt=0.32t(1t)\gamma_t=0.3\sqrt{2t(1-t)} 的模型,兩個 head 同時輸出 btb_tηt\eta_t。取樣器:

def sde_step(x, t, h, eps_t):
    b, eta = model(x, t)
    s = -eta / gamma(t)
    return x + h * (b + eps_t * s) + (2 * eps_t * h) ** 0.5 * torch.randn_like(x)

ε{0,0.1,0.5}\varepsilon\in\lbrace 0,\,0.1,\,0.5 \rbrace 與步數 {10,50,200,1000}\lbrace 10,50,200,1000 \rbraceW2W_2,畫成三條曲線。先寫下你的猜測:這三條會不會交叉? 這張圖要驗的是 U1.4 的判準——不是「哪一個比較好」,而是「在這個設定下,哪一種誤差佔上風」。

補充這一步很容易得到和文獻相反的結論,原因在哪

「少步數 ODE 贏、多步數 SDE 贏」這個說法有前提,而那些前提很容易在 toy 上不成立。

第一,εt\varepsilon_t 的尺度。如果照 εt=cγt2\varepsilon_t=c\,\gamma_t^2 取而 γ\gamma 又很小,那每步注入的噪聲 2εh\sqrt{2\varepsilon h} 就很小,而 score 那一項提供的是一個很強的收縮——這時候 SDE 幾乎在所有步數都比 ODE 好,因為它其實是一個更穩定的積分器,不是「多了噪聲」。要重現文獻的圖像,ε\varepsilon 得取到 diffusion reverse SDE 那個量級。

第二,score 是精確的還是學出來的。SDE 那一族的邊際不變性用到了 sts_t 是真正的 score(U3.1 的證明只有那一步用到它)。sts_t 有誤差時,ε\varepsilon 越大偏得越多——文獻上「多步 SDE 更好」的圖像其實同時包含了「噪聲能抵掉一部分網路誤差」,而 toy 上如果 score 幾乎精確,就看不到這一面。

所以這一步的正確做法是:先把 ε\varepsilon 掃過一個大範圍(不要只試一兩個值),並且同時報告用精確 score 與用訓練 score 的兩組曲線。得到和文獻不同的結論不一定是做錯了,但一定要說清楚是在哪個設定下量到的。

再做一件事:在 t=0.3t=0.3 對所有樣本注入 +0.5+0.5 的人為偏移,看三條曲線終點的 W2W_2ε=0\varepsilon=0 的偏移會完整保留,ε>0\varepsilon \gt 0 會被吸收掉一部分——U3.1 的 demo 量過,吸收的程度取決於 ε\varepsilon 的大小和剩下多少時間。

課堂提問Q1

圖 b 上有一個點明顯落在那條共線之上——同樣的曲率積分,誤差卻大了一截。

第一個該懷疑什麼?

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

不是曲率積分算錯,也不是步數不夠(步數只會讓整條線平移)。

回到 U3.2 的不等式:

Euler 全域誤差  CLh01x¨tdt.\text{Euler 全域誤差}\ \lesssim\ C_L\,h\int_0^1\|\ddot x_t\|\,dt .

它有兩個因子。畫散點圖的時候我們默默假設了各個模型的 CLC_L 差不多,於是誤差只隨積分變化。所以某個點偏高,第一個該懷疑的就是那個模型的速度場 Lipschitz 常數比較大——場更「陡」,同樣的局部誤差被放大得更多。

怎麼檢查?幾個實際的做法:

  • 在一批隨機的 (x,t)(x,t) 上估 uθ/x\|\partial u_\theta/\partial x\| 的譜範數(用幾次 power iteration,或直接 finite difference 取最大方向)。
  • 看那個模型的權重範數與有沒有用 weight decay/spectral norm。
  • 看它的誤差是不是集中在少數幾個起點(CLC_L 大通常表現成「大部分軌跡很好,少數幾條爆掉」),而不是均勻分布。

如果 CLC_L 也差不多,那才輪到懷疑量測本身:曲率積分用的 n_fine 夠不夠細(差分兩次對步長很敏感)、W2W_2 的樣本數夠不夠、以及有沒有樣本走到不同的 mode(那種樣本的誤差是 O(1)O(1),和步數幾乎無關,會把平均值整個拉歪——U2.5 用中位數就是為了避開它)。

點偏離那條線的時候,先懷疑 CLC_L,不要先懷疑理論。

作業

  1. Guidance 與曲率。 用兩類月牙訓一個條件 FM(條件 dropout 0.1)。對 w{1,2,4,8}w\in\lbrace 1,2,4,8 \rbrace 算曲率積分與 Euler 10 步的 W2W_2。它們還落在圖 b 的那條線上嗎?順便量一下終點的散布,對照 U3.5 demo 的結論。
  2. Heun。 用 Heun 重跑步驟 3。斜率變成多少?哪個模型改善最多、哪個最少?用 U3.5 的三階導數論證解釋。
  3. 時間加權。 回到 U2.5 作業 2:用 logit-normal 抽 tt 重訓獨立配對的 FM,畫它在圖 b 上的位置。它移動的是橫座標還是縱座標?為什麼?
  4. γ\gamma 的張力。γ\gamma 幅度 {0,0.1,0.3,1.0}\lbrace 0,\,0.1,\,0.3,\,1.0 \rbrace 各訓一個 OT 配對的模型,畫「曲率積分」與「最佳 ε\varepsilon 下 SDE 1000 步的 W2W_2」對 γ\gamma 的兩條曲線。這就是 U3.2 Q1 那個張力的實測——注意 U3.2 demo 量到 γ\gamma 的效果取決於配對好不好,先猜再跑。
  5. (選)高維。 在 MNIST 上重做步驟 1 的 (a) 與 (d),比較 Euler 10 步的 FID。增益是不是像 U3.4 預測的比 2D 小?順便量一下成本矩陣的變異係數,看它是不是接近 2/d\sqrt{2/d}

消化一下

想一想

圖 b 中若某個模型的點明顯落在共線之上(同樣曲率、誤差更大),最可能的原因是什麼?

想一想

步驟 5 中若 ε>0\varepsilon \gt 0 的曲線在少步時較差,理由是什麼?

想一想

根據 demo 量到的數字,關於「直」與「省」下列哪一句正確?

參考文獻

  1. Liu, X., Gong, C., Liu, Q. Flow Straight and Fast: Learning to Generate and Transfer Data with Rectified Flow. ICLR 2023. (步驟 1 的 (b)(c)。)
  2. Tong, A., et al. Improving and Generalizing Flow-Based Generative Models with Minibatch Optimal Transport. TMLR 2024. (步驟 1 的 (d)。)
  3. Albergo, M. S., Boffi, N. M., Vanden-Eijnden, E. Stochastic Interpolants: A Unifying Framework for Flows and Diffusions. 2023. (步驟 5 的 γt\gamma_t 訓練與一族取樣器。)