U3.4 拉直(二):在每個 batch 裡解一次配對
本篇重用M6.0搬倉庫問題:Monge、Kantorovich 與 Coupling·M6.2最省油的配法一定不交叉:一維的單調性與 Brenier·M6.3五萬個倉庫怎麼算:Sinkhorn 與 Minibatch 近似
如果可以直接挑最好的配對呢
U3.3 的做法是繞一圈:先訓一個模型、用它取樣、把結果當新配對。但 U3.2 的結論其實更直接——我們要的是一組不交叉的配對。那能不能不繞,一開始就挑一組好的?
「好」在這裡有一個現成的定義:讓總位移最小,也就是 U1.0 倉庫搬貨那個問題的最佳解——optimal transport。
最省力的那組配對,
是不是剛好也是不交叉的那一組?
交換一次就看出來了
是。而且理由用一行計算就講完。
假設某個配對把 送到 、把 送到 。如果交換終點(、),二次成本的變化是
(展開之後 全部消掉,只剩交叉項。)
所以:如果 ,交換一定更省。 一個最佳的配對不可能還有這種可以改進的對,於是它必須滿足
這就是單調。在一維,單調直接等於「位移直線完全不交叉」(兩條相交的線交換之後總長更短)。在高維,單調不是「零交叉」的等價敘述,但它把交叉壓到很少——下面 demo 量到的是從 20% 掉到 1% 以下。
補充Brenier 定理說了什麼
在 有密度、二次成本的條件下,OT 耦合其實是一個確定映射 ,而且 是某個凸函數的梯度。凸函數的梯度天生單調,這是上面那個不等式的另一種來源。
值得注意的是這裡出現了 U3.0 講過的張力:OT 配對是確定性的,所以 時沒有可回歸的 score。不過minibatch OT 沒有這個問題,理由可以說得很具體:同一個 在不同的 batch 裡會遇到不同的 32 個候選,解出來的最佳配對也就不同。所以 不是一個函數,而是一個隨 batch 變動的隨機指派——給定 仍然有後驗可以平均。batch 越大越接近全域 OT,那個隨機性就越小; 的極限才真的塌成確定映射。
於是 OT 配對之下的 flow matching 就是最理想的情況:條件直線幾乎不交叉、邊際速度幾乎等於條件速度、、一步 Euler 幾乎精確。它正是 U3.3 迭代想逼近的那個極限。
問題只有一個:全域 OT 要在 個資料點上解一個 的指派問題。 是五萬張圖的時候,這件事做不到。
Minibatch OT:只多兩行
Tong 等人的 OT-CFM [1] 與 Pooladian 等人的 multisample FM [2] 的做法乾脆得有點好笑:只在每個 batch 裡解 OT。
repeat
x1 ← B 個資料樣本
x0 ← B 個噪聲樣本
C ← 成本矩陣 C[i,j] = ‖x0[i] − x1[j]‖²
σ ← 解指派問題 argmin_σ Σ_i C[i, σ(i)] # scipy.optimize.linear_sum_assignment
x1 ← x1[σ] # 重排,讓 x0[i] 對到 x1[σ(i)]
之後與 conditional flow matching 完全相同:抽 t、算 xt、target = x1 − x0
多的就是那兩行。精確指派是 , 時和一次前向傳播比起來可以忽略;batch 再大就換 Sinkhorn 近似 [3]。
互動 demo:batch 內解一次 OT。 從「獨立」按到「192(全域)」,紅色那些「和別人交叉的配對」一路減少。實測交叉比例約 20% → 9% → 4% → 1.4% → 0.6%,1 步 Euler 的誤差 1.5 → 0.77 → 0.41 → 0.21 → 0.08。batch 16 就走完八成的路,但後面還有得賺——是報酬遞減,不是「16 就夠了」。
OT 配對真正保證的是「單調」,不是「不交叉」。
一維的單調等於零交叉,高維只能說把交叉壓到很少。
邊際正確,配對次佳
Batch 內重排只是一個置換: 的集合和 的集合都沒有變,只有誰配誰改了。所以 minibatch OT 給出的 仍然是 與 的合法耦合,U2.2 那兩個定理照樣成立,訓出來的速度場生成的終點分布仍然是 。沒有邊際偏差, 多小都沒有。
偏差在別的地方: 不是全域 OT。一個 batch 裡的 只能配到同一個 batch 裡的 ,而它真正的 OT 對象大概率不在這個 batch 裡。所以 是「局部最佳、全域次佳」,條件直線還是會交叉(只是比獨立配對少很多),。 時 才趨近全域 OT。
Minibatch OT 影響的是曲率,不是正確性。
講精確一點:它把你在 U3.2 那張散點圖上的位置往左移(曲率積分變小),而不是改變那條線的斜率。斜率是取樣器的階數,換配對動不了它。
這一點和 reflow 不同:reflow 第 輪的配對帶著第 輪模型的誤差,所以邊際有可能偏。
reflow 與 minibatch OT,差在哪?
| Reflow | Minibatch OT | |
|---|---|---|
| 改配對的時機 | 訓完一輪之後 | 每個 batch 當下 |
| 額外計算 | 整輪 ODE 取樣 + 重訓 | 每 batch 一次 指派 |
| 邊際 | 會累積模型誤差 | 精確(只是置換) |
| 極限 | 直(),不一定是 OT | 全域 OT() |
| 配對是否確定性 | 是(所以沒有可回歸的 score) | 不是(整體是很多置換的混合) |
| 高維表現 | 不受距離集中影響 | 成本矩陣退化,增益減弱 |
課堂提問Q1
最後一列是實務上最重要的一格。在 2D toy 上 minibatch OT 的效果非常漂亮,但在圖像( 以上)上增益小得多。
為什麼?
先想一想,再展開看整理後的答案
因為高維的距離會集中。
取兩個獨立的 ,則
所以相對標準差是 —— 越大,所有兩兩距離越接近同一個值。 時是 100%, 時只剩 2.5%。
成本矩陣幾乎是常數的時候,「最佳指派」和「隨機指派」的總成本差不了多少,於是 OT 配對能買到的結構就很少。這不是實作沒調好,而是幾何本身的性質。
兩個實務上的推論:
- 在 latent space 做比在 pixel space 做有效。 latent 的維度低(例如 壓到幾千維以下),而且不同 latent 之間的距離結構比 pixel 有意義。
- batch size 要大才看得到效果。 小的時候一個 batch 內的點太少,本來就配不到好對象;而高維又讓「好對象」和「壞對象」的差距變小,兩個效應疊起來。
順帶一提,這也解釋了為什麼 U3.3 在高維反而比較常見:reflow 用的是模型自己的 ODE 誘導的配對,那個配對的結構是學出來的,不依賴成本矩陣有沒有對比度。
消化一下
參考文獻
- Tong, A., Fatras, K., Malkin, N., Huguet, G., Zhang, Y., Rector-Brooks, J., Wolf, G., Bengio, Y. Improving and Generalizing Flow-Based Generative Models with Minibatch Optimal Transport. TMLR 2024. (OT-CFM;這一篇的主線。)
- Pooladian, A.-A., Ben-Hamu, H., Domingo-Enrich, C., Amos, B., Lipman, Y., Chen, R. T. Q. Multisample Flow Matching: Straightening Flows with Minibatch Couplings. ICML 2023. (同一個想法的另一條線,含 batch size 的分析。)
- Cuturi, M. Sinkhorn Distances: Lightspeed Computation of Optimal Transport. NeurIPS 2013. (batch 太大時用的熵正則化近似。)