U3.6 實作:把曲率積分畫成一條線
本篇重用M2.3導航每 30 秒更新一次:Euler 法與它的誤差·M6.1兩個城市的人口分佈差多少:W₂ 距離
設定
沿用 toy.py 與 fm.py。這個單元新增三個檔:couplings.py(reflow 的配對生成、minibatch OT)、interpolant.py(帶 的訓練與一整族 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 積分、再對起點平均
對四個模型算 ,以及 Euler 10 步終點的 ,畫成散點、兩軸取對數。
會看到的事情是:四個點近似共線、斜率約 1。這張圖就是 U3.2 那條不等式的實證——高度由曲率決定。U3.2 的 demo 已經先用精確場把這條線畫出來過(六個設定、 通常 0.95 以上),這裡要看的是換成訓練出來的網路之後還成不成立。
再加一組步數():點會整體上下平移,但相對順序不變。
曲率積分小,誤差就一定小嗎?
這件事能不能在自己的圖上驗一次?
步驟 4:straightness 與運輸成本
對每個模型算 (U3.3)與 ,做成一張表。
會看到的事情是:reflow 的兩個數都單調下降,而 OT 的運輸成本最低。把 OT 的成本和全域 OT(2000 點精確解)比一下,就知道 reflow 的極限離 OT 有多遠。
步驟 5:ODE 與 SDE 在第幾步交叉(圖 c)
用 interpolant.py 訓一個 的模型,兩個 head 同時輸出 與 。取樣器:
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)
對 與步數 算 ,畫成三條曲線。先寫下你的猜測:這三條會不會交叉? 這張圖要驗的是 U1.4 的判準——不是「哪一個比較好」,而是「在這個設定下,哪一種誤差佔上風」。
補充這一步很容易得到和文獻相反的結論,原因在哪
「少步數 ODE 贏、多步數 SDE 贏」這個說法有前提,而那些前提很容易在 toy 上不成立。
第一, 的尺度。如果照 取而 又很小,那每步注入的噪聲 就很小,而 score 那一項提供的是一個很強的收縮——這時候 SDE 幾乎在所有步數都比 ODE 好,因為它其實是一個更穩定的積分器,不是「多了噪聲」。要重現文獻的圖像, 得取到 diffusion reverse SDE 那個量級。
第二,score 是精確的還是學出來的。SDE 那一族的邊際不變性用到了 是真正的 score(U3.1 的證明只有那一步用到它)。 有誤差時, 越大偏得越多——文獻上「多步 SDE 更好」的圖像其實同時包含了「噪聲能抵掉一部分網路誤差」,而 toy 上如果 score 幾乎精確,就看不到這一面。
所以這一步的正確做法是:先把 掃過一個大範圍(不要只試一兩個值),並且同時報告用精確 score 與用訓練 score 的兩組曲線。得到和文獻不同的結論不一定是做錯了,但一定要說清楚是在哪個設定下量到的。
再做一件事:在 對所有樣本注入 的人為偏移,看三條曲線終點的 。 的偏移會完整保留, 會被吸收掉一部分——U3.1 的 demo 量過,吸收的程度取決於 的大小和剩下多少時間。
課堂提問Q1
圖 b 上有一個點明顯落在那條共線之上——同樣的曲率積分,誤差卻大了一截。
第一個該懷疑什麼?
先想一想,再展開看整理後的答案
不是曲率積分算錯,也不是步數不夠(步數只會讓整條線平移)。
回到 U3.2 的不等式:
它有兩個因子。畫散點圖的時候我們默默假設了各個模型的 差不多,於是誤差只隨積分變化。所以某個點偏高,第一個該懷疑的就是那個模型的速度場 Lipschitz 常數比較大——場更「陡」,同樣的局部誤差被放大得更多。
怎麼檢查?幾個實際的做法:
- 在一批隨機的 上估 的譜範數(用幾次 power iteration,或直接 finite difference 取最大方向)。
- 看那個模型的權重範數與有沒有用 weight decay/spectral norm。
- 看它的誤差是不是集中在少數幾個起點( 大通常表現成「大部分軌跡很好,少數幾條爆掉」),而不是均勻分布。
如果 也差不多,那才輪到懷疑量測本身:曲率積分用的 n_fine 夠不夠細(差分兩次對步長很敏感)、 的樣本數夠不夠、以及有沒有樣本走到不同的 mode(那種樣本的誤差是 ,和步數幾乎無關,會把平均值整個拉歪——U2.5 用中位數就是為了避開它)。
點偏離那條線的時候,先懷疑 ,不要先懷疑理論。
作業
- Guidance 與曲率。 用兩類月牙訓一個條件 FM(條件 dropout 0.1)。對 算曲率積分與 Euler 10 步的 。它們還落在圖 b 的那條線上嗎?順便量一下終點的散布,對照 U3.5 demo 的結論。
- Heun。 用 Heun 重跑步驟 3。斜率變成多少?哪個模型改善最多、哪個最少?用 U3.5 的三階導數論證解釋。
- 時間加權。 回到 U2.5 作業 2:用 logit-normal 抽 重訓獨立配對的 FM,畫它在圖 b 上的位置。它移動的是橫座標還是縱座標?為什麼?
- 的張力。 對 幅度 各訓一個 OT 配對的模型,畫「曲率積分」與「最佳 下 SDE 1000 步的 」對 的兩條曲線。這就是 U3.2 Q1 那個張力的實測——注意 U3.2 demo 量到 的效果取決於配對好不好,先猜再跑。
- (選)高維。 在 MNIST 上重做步驟 1 的 (a) 與 (d),比較 Euler 10 步的 FID。增益是不是像 U3.4 預測的比 2D 小?順便量一下成本矩陣的變異係數,看它是不是接近 。
消化一下
參考文獻
- Liu, X., Gong, C., Liu, Q. Flow Straight and Fast: Learning to Generate and Transfer Data with Rectified Flow. ICLR 2023. (步驟 1 的 (b)(c)。)
- Tong, A., et al. Improving and Generalizing Flow-Based Generative Models with Minibatch Optimal Transport. TMLR 2024. (步驟 1 的 (d)。)
- Albergo, M. S., Boffi, N. M., Vanden-Eijnden, E. Stochastic Interpolants: A Unifying Framework for Flows and Diffusions. 2023. (步驟 5 的 訓練與一族取樣器。)