M6.3 五萬個倉庫怎麼算:Sinkhorn 與 Minibatch 近似
本篇重用M6.0搬倉庫問題:Monge、Kantorovich 與 Coupling·M6.1兩個城市的人口分佈差多少:W₂ 距離·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy
二十五億格的表
M6.0 的運輸表在三個倉庫、三家店時是九格,LP 一瞬間解完。現在倉庫與店各有五萬個:表有 格。光是把成本矩陣存進記憶體就要二十 GB,LP 的複雜度更是 量級——算不動。
而「五萬個」在現實裡很常見:五萬張圖片的特徵、五萬個顧客、五萬個粒子。你需要 或最佳配對,但不可能精確算。
先用直覺想兩條路,並寫下你覺得各自會偏在哪個方向:(一)不要求「最省」,接受一個「差不多省、但好算」的搬法;(二)不要一次配五萬個,隨機抽兩百個倉庫、兩百家店,在小組裡配,多抽幾次取平均。第一條路的答案會比真實成本高還是低?第二條路算出來的距離會比真實 大還是小?
課堂提問Q1
翻成數學。「差不多省、但好算」可以怎麼寫成一個修改過的最佳化問題?「在小組裡配」估的是什麼量——它和 是同一個東西嗎?
先想一想,再展開看整理後的答案
課堂上第一條路的候選很多(貪婪配對、只看最近的幾家店…),但有一個修改特別乾淨:在成本上加一個懲罰「太集中」的項。運輸表 越集中(每列只有一格不為零)越像精確解、越難算;越分散(接近獨立 coupling)越平滑。用 entropy 量「分散程度」,解
這叫 entropic optimal transport。 是「願意為了好算而多付多少油錢」的兌換率。它的最佳解會比精確解分散、成本偏高——這是第一條路的偏差方向,下面會證明。
第二條路:抽 個倉庫、 家店,在小組裡精確解 ( 小到可以用 LP 或排序),對很多組取平均:, 是 個樣本的經驗分佈。它不是 :一個小組裡的兩百個點根本代表不了整個分佈,硬要在組內把每個倉庫配到某家店,會被迫配到不該配的地方。這條路的偏差方向是:估出來的距離高於 ,配對比真實最佳配對更亂。
先抓住這個畫面:把稜角磨圓,或只看一小塊
精確 OT 的解是運輸多面體的一個頂點——一個有稜角的東西,最多 格不為零。要找到頂點得走 LP 那種組合式的路。
Entropic OT 把稜角磨圓。 加了 之後,目標是 strictly convex 的,最佳解在多面體內部、每一格都嚴格大於零,而且有一個極簡的形狀:一個固定的矩陣 ,左右各乘一個對角縮放。要找到它,只要輪流把列和調成 、把欄和調成 ——每一輪是一次矩陣乘向量。這就是 Sinkhorn。
Minibatch 只看一小塊。 五萬對五萬看不完,就看兩百對兩百。組內的問題小到能精確解。代價是每一組都在「以偏概全」:組裡剛好沒有某個區域的店,那個區域的倉庫就被迫送遠。
寫成數學:Sinkhorn
解的形狀。 對 entropic 問題寫 Lagrangian,對每個 微分令其為零:
是列和、欄和兩組約束的 Lagrange multiplier。所以最佳解一定是 : 由成本與 決定、只算一次;未知的只有 個縮放因子 ,不是 個格子。
Sinkhorn 演算法。 把 調到列和是 、欄和是 。兩個約束分開看各自是一行:
輪流做——固定 解 、固定 解 ——每一步都精確滿足其中一組約束、稍微破壞另一組。Sinkhorn & Knopp [3] 證明它收斂到唯一的 (差一個常數倍);Cuturi [1] 把它帶進大規模 OT。每一輪是兩次矩陣乘向量,;五萬對五萬時 仍是 格,但不需要 LP 的組合搜尋,而且矩陣乘向量可以平行、可以用 GPU,幾十輪就收斂。
兩端的極限。 : 全 1 的矩陣, 獨立 coupling ——完全不看成本。: 消失,回到精確 OT(但 的元素會下溢到 0,數值上要在 域做——展開框)。中間的 是取捨:越小越準、越大越快收斂、越穩定。
偏差的方向。 精確解 的成本 是所有 coupling 裡最小的;entropic 解 是另一個 coupling,所以 ——entropic 解的運輸成本永遠不低於真實 OT。它多付的成本換來的是 :拆得更散。
展開細節log 域的 Sinkhorn、Sinkhorn divergence、以及 ε 的偏差量級
數值上 K = e^{−c/ε} 在 ε 小時會下溢(c/ε ≈ 700 就變成 0)。改在 log 域做:令 f = ε log u、g = ε log v,更新變成 f_i = ε log a_i − ε log Σ_j exp((g_j − c_{ij})/ε),用 log-sum-exp 穩定計算。這是實務上的標準寫法。
entropic 成本 ⟨π_ε, c⟩ 與精確 OT 的差是 O(ε log(1/ε)) 量級(Peyré & Cuturi [2] 第 4.1 節);ε → 0 時收斂但很慢。另外 entropic OT 有一個討厭的性質:把同一個分佈搬到自己,成本不是零(π_ε 會把質量拆散到鄰近的點)。Genevay et al. [5] 的修正是 Sinkhorn divergence:S_ε(μ,ν) = OT_ε(μ,ν) − ½OT_ε(μ,μ) − ½OT_ε(ν,ν),把「搬到自己也有成本」的那部分扣掉;它在 μ = ν 時為零、對 ε 的偏差小得多,是拿 entropic OT 當「距離」用時的標準版本。
寫成數學:Minibatch 的偏差
抽 個 的樣本、 個 的樣本,各成經驗分佈 ,算 ,對抽樣取期望。兩件事可以說清楚。
距離偏高。 這可以證明: 對 這一對是jointly convex 的(兩個 coupling 的凸組合仍是對應邊際凸組合的 coupling,所以最小值不會高於凸組合),而 、,M5.0 立刻給
直覺是:組內的配對被「組裡只有這 家店」的限制綁住,每個倉庫的最近選項比全體裡少,被迫送得更遠。極端例子:,真實 ;但兩組獨立抽出的 個點幾乎不會重合, 嚴格成立,且在 維以 的速率緩慢趨近零(經驗分佈的 收斂速率,見 Peyré & Cuturi [2])。所以 minibatch 對「兩個分佈其實相同」的判斷有系統性偏差, 越小、維度越高越嚴重 [4]。
配對偏亂。 把每一組內的最佳配對合起來,得到一個「minibatch coupling」。它是 的合法 coupling(每組內邊際都對,平均後也對),所以它的成本 ——又是「任何 coupling 給上界」。而且它不是一個映射:同一個倉庫在不同組裡會被配到不同的店,平均後一個 對應到一片 。 時 , 時 (獨立 coupling)—— 在「精確 OT」與「完全隨機配」之間插值,角色與 Sinkhorn 的 相同。
兩種近似殊途同歸:都給出一個比精確解更分散的 coupling、成本偏高;一個用 控制、一個用 控制。
回到五萬個倉庫:該用哪個
回答起點問題。第一條路(entropic)成本偏高、配對偏散,偏差由 控制,代價是 的矩陣乘向量——五萬對五萬在 GPU 上可行,但 要 格的記憶體(可以分塊)。第二條路(minibatch)成本偏高、配對偏亂,偏差由 控制,記憶體只要 ,但每組的估計有隨機性、要多抽幾組平均。
選擇的判準:你要的是「距離」還是「配對」? 若要的是一個數字(兩個大點雲差多少),Sinkhorn divergence 加中等的 通常最好——偏差可控、無隨機性。若要的是「每個 該配哪個 」而且要在訓練迴圈裡每一步都算,minibatch 是唯一現實的選擇——這時你要清楚配對是「偏亂」的, 越大越接近真實 OT。
依賴的假設:Sinkhorn 的收斂需要 沒有全零的列或欄( 太小時會發生,要用 log 域);minibatch 的偏差分析假設樣本獨立同分佈抽自 。兩者的偏差方向(成本偏高)是嚴格的,但偏差的大小要靠實驗量。
回到五萬個倉庫:縮小成兩千個,把三個數並排
import numpy as np
from scipy.special import logsumexp
from scipy.optimize import linear_sum_assignment
rng = np.random.default_rng(0)
n, d = 2000, 2
X = rng.normal(0, 1, (n, d)); Y = rng.normal(0, 1, (n, d)) @ np.diag([2, 0.5]) + np.array([3, 0])
print(9 + (1-2)**2 + (1-0.5)**2) # closed form W₂² = 10.25(μ=N(0,I)、ν=N((3,0), diag(4,¼)))
C = ((X[:, None, :] - Y[None, :, :])**2).sum(-1) # 2000×2000 成本矩陣
r, c = linear_sum_assignment(C); print(C[r, c].mean()) # 2000 個樣本的精確 OT ≈ 10.73(抽樣誤差在上方)
def sinkhorn_log(C, eps, iters=1000): # log 域的 Sinkhorn:輪流修列和、欄和
m = C.shape[0]; loga = -np.log(m) * np.ones(m); f = np.zeros(m); g = np.zeros(m)
for _ in range(iters):
f = eps * (loga - logsumexp((g[None, :] - C) / eps, axis=1)) # 修列和
g = eps * (loga - logsumexp((f[:, None] - C) / eps, axis=0)) # 修欄和
logP = (f[:, None] + g[None, :] - C) / eps; P = np.exp(logP)
return (P * C).sum(), -(P * logP).sum()
for eps in [0.05, 0.2, 1.0, 5.0]:
cost, H = sinkhorn_log(C, eps); print(eps, round(cost, 2), round(H, 2)) # 成本 10.77 → 13.31 單調上升;entropy 上升
def minibatch_w2(C, B, reps=200): # 小組內用 assignment 精確解
vals = []
for _ in range(reps):
i = rng.choice(n, B, replace=False); j = rng.choice(n, B, replace=False)
S = C[np.ix_(i, j)]; r, c = linear_sum_assignment(S); vals.append(S[r, c].mean())
return np.mean(vals)
for B in [8, 32, 128, 512]: print(B, round(minibatch_w2(C, B), 2)) # 11.91 → 10.77:從上方逼近 10.73
# 破壞:μ = ν(真值 0)
Y2 = rng.normal(0, 1, (n, d)); C2 = ((X[:, None, :] - Y2[None, :, :])**2).sum(-1)
for B in [8, 512]: print('same', B, round(minibatch_w2(C2, B, 50), 3)) # 1.33、0.064:相同分佈被判成有差
三組數字並排。closed form 真值 ;用全部 2000 個樣本精確解 assignment 得 (有限樣本的估計本身就偏高一點)。Sinkhorn 成本隨 從 單調升到 、entropy 從 升到 ——偏高、偏散,而且 就已經在 下溢的邊緣,所以程式直接寫成 域。minibatch 從 的 降到 的 ,從上方逼近 。
這裡有一個實作上很會騙人的地方,值得自己踩一次。把 iters 從 1000 降回 300, 會給出 ——比精確解 還低,直接違反上面那個「entropic 解的成本不低於真實 OT」。原因不是定理錯了,是那個 還不是合法的 coupling:300 輪之後列邊際的 誤差還有 ,質量沒有守好,成本自然可以被壓到下界以下。 越小收斂越慢,所以看到「比精確解還低」的第一件事是去檢查邊際,不是去懷疑定理。
刻意違反一個假設。 把 換成與 同分佈的另一組樣本(真值 ):2000 對 2000 的精確解給 ,minibatch 在 時給 、 時給 ——它會把「兩個相同的分佈」判成「有差」, 越小越嚴重。再把 壓到 且改回普通域的 :矩陣大量下溢成 0、K @ v 出現零、除法變成 inf——這就是為什麼要在 域做。
先消化一下
參考文獻
- Cuturi, M. Sinkhorn Distances: Lightspeed Computation of Optimal Transport. NeurIPS 2013.(把 entropic 正則化與 Sinkhorn 迭代帶進大規模 OT。)
- Peyré, G., Cuturi, M. Computational Optimal Transport. Foundations and Trends in Machine Learning 11(5–6), 2019.(第 4 章:entropic OT、Sinkhorn 的推導與收斂、log 域實作、 的偏差量級;第 8.5 節:經驗分佈的收斂速率。)
- Sinkhorn, R., Knopp, P. Concerning Nonnegative Matrices and Doubly Stochastic Matrices. Pacific Journal of Mathematics 21(2), 1967.(交替正規化收斂到唯一的對角縮放。)
- Fatras, K., Zine, Y., Flamary, R., Gribonval, R., Courty, N. Learning with Minibatch Wasserstein: Asymptotic and Gradient Properties. AISTATS 2020.(minibatch OT 的偏差、 時不為零、隨 的收斂。)
- Genevay, A., Peyré, G., Cuturi, M. Learning Generative Models with Sinkhorn Divergences. AISTATS 2018.(Sinkhorn divergence:扣掉自我運輸成本。)