M6.3 18 分鐘閱讀 2026年9月

M6.3 五萬個倉庫怎麼算:Sinkhorn 與 Minibatch 近似

本篇重用M6.0搬倉庫問題:Monge、Kantorovich 與 Coupling·M6.1兩個城市的人口分佈差多少:W₂ 距離·M5.1用錯的機率表下注一年會多輸多少:KL Divergence 與 Cross-Entropy

二十五億格的表

M6.0 的運輸表在三個倉庫、三家店時是九格,LP 一瞬間解完。現在倉庫與店各有五萬個:表有 5×104×5×104=2.5×1095\times10^4\times5\times10^4=2.5\times10^9 格。光是把成本矩陣存進記憶體就要二十 GB,LP 的複雜度更是 n3n^3 量級——算不動。

而「五萬個」在現實裡很常見:五萬張圖片的特徵、五萬個顧客、五萬個粒子。你需要 W2W_2 或最佳配對,但不可能精確算。

先用直覺想兩條路,並寫下你覺得各自會偏在哪個方向:(一)不要求「最省」,接受一個「差不多省、但好算」的搬法;(二)不要一次配五萬個,隨機抽兩百個倉庫、兩百家店,在小組裡配,多抽幾次取平均。第一條路的答案會比真實成本高還是低?第二條路算出來的距離會比真實 W2W_2 大還是小?

課堂提問Q1

翻成數學。「差不多省、但好算」可以怎麼寫成一個修改過的最佳化問題?「在小組裡配」估的是什麼量——它和 W2(μ,ν)W_2(\mu,\nu) 是同一個東西嗎?

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

課堂上第一條路的候選很多(貪婪配對、只看最近的幾家店…),但有一個修改特別乾淨:在成本上加一個懲罰「太集中」的項。運輸表 π\pi 越集中(每列只有一格不為零)越像精確解、越難算;越分散(接近獨立 coupling)越平滑。用 entropy H(π)=πijlogπijH(\pi)=-\sum\pi_{ij}\log\pi_{ij} 量「分散程度」,解

minπΠ(a,b) ijπijcij    εH(π),ε>0.\min_{\pi\in\Pi(a,b)}\ \sum_{ij}\pi_{ij}c_{ij}\;-\;\varepsilon\,H(\pi),\qquad\varepsilon>0 .

這叫 entropic optimal transportε\varepsilon 是「願意為了好算而多付多少油錢」的兌換率。它的最佳解會比精確解分散、成本偏高——這是第一條路的偏差方向,下面會證明。

第二條路:抽 BB 個倉庫、BB 家店,在小組裡精確解 W2W_2BB 小到可以用 LP 或排序),對很多組取平均:E[W2(μ^B,ν^B)]\mathbb E\big[W_2(\hat\mu_B,\hat\nu_B)\big]μ^B\hat\mu_BBB 個樣本的經驗分佈。它不是 W2(μ,ν)W_2(\mu,\nu):一個小組裡的兩百個點根本代表不了整個分佈,硬要在組內把每個倉庫配到某家店,會被迫配到不該配的地方。這條路的偏差方向是:估出來的距離高於 W2(μ,ν)W_2(\mu,\nu),配對比真實最佳配對更亂

先抓住這個畫面:把稜角磨圓,或只看一小塊

精確 OT 的解是運輸多面體的一個頂點——一個有稜角的東西,最多 m+n1m+n-1 格不為零。要找到頂點得走 LP 那種組合式的路。

Entropic OT 把稜角磨圓。 加了 εH(π)-\varepsilon H(\pi) 之後,目標是 strictly convex 的,最佳解在多面體內部、每一格都嚴格大於零,而且有一個極簡的形狀:一個固定的矩陣 Kij=ecij/εK_{ij}=e^{-c_{ij}/\varepsilon},左右各乘一個對角縮放。要找到它,只要輪流把列和調成 aa、把欄和調成 bb——每一輪是一次矩陣乘向量。這就是 Sinkhorn。

Minibatch 只看一小塊。 五萬對五萬看不完,就看兩百對兩百。組內的問題小到能精確解。代價是每一組都在「以偏概全」:組裡剛好沒有某個區域的店,那個區域的倉庫就被迫送遠。

寫成數學:Sinkhorn

解的形狀。 對 entropic 問題寫 Lagrangian,對每個 πij\pi_{ij} 微分令其為零:

cij+ε(logπij+1)figj=0πij=e(fiε)/εui  ecij/εKij  egj/εvj,c_{ij}+\varepsilon(\log\pi_{ij}+1)-f_i-g_j=0 \quad\Longrightarrow\quad \pi_{ij}=\underbrace{e^{(f_i-\varepsilon)/\varepsilon}}_{u_i}\;\underbrace{e^{-c_{ij}/\varepsilon}}_{K_{ij}}\;\underbrace{e^{g_j/\varepsilon}}_{v_j},

f,gf,g 是列和、欄和兩組約束的 Lagrange multiplier。所以最佳解一定是 π=diag(u)Kdiag(v)\pi=\mathrm{diag}(u)\,K\,\mathrm{diag}(v)KK 由成本與 ε\varepsilon 決定、只算一次;未知的只有 m+nm+n 個縮放因子 u,vu,v,不是 mnmn 個格子。

Sinkhorn 演算法。u,vu,v 調到列和是 aa、欄和是 bb。兩個約束分開看各自是一行:

列和=a:  ui=ai(Kv)i,欄和=b:  vj=bj(K ⁣u)j.\text{列和}=a:\ \ u_i=\frac{a_i}{(Kv)_i},\qquad \text{欄和}=b:\ \ v_j=\frac{b_j}{(K^{\!\top}u)_j}.

輪流做——固定 vvuu、固定 uuvv——每一步都精確滿足其中一組約束、稍微破壞另一組。Sinkhorn & Knopp [3] 證明它收斂到唯一的 (u,v)(u,v)(差一個常數倍);Cuturi [1] 把它帶進大規模 OT。每一輪是兩次矩陣乘向量,O(mn)O(mn);五萬對五萬時 KK 仍是 2.5×1092.5\times10^9 格,但不需要 LP 的組合搜尋,而且矩陣乘向量可以平行、可以用 GPU,幾十輪就收斂。

ε\varepsilon 兩端的極限。 ε\varepsilon\to\inftyKK\to 全 1 的矩陣,π\pi\to 獨立 coupling abab^\top——完全不看成本。ε0\varepsilon\to0εH-\varepsilon H 消失,回到精確 OT(但 K=ec/εK=e^{-c/\varepsilon} 的元素會下溢到 0,數值上要在 log\log 域做——展開框)。中間的 ε\varepsilon 是取捨:越小越準、越大越快收斂、越穩定。

偏差的方向。 精確解 π\pi^* 的成本 π,c\langle\pi^*,c\rangle 是所有 coupling 裡最小的;entropic 解 πε\pi_\varepsilon 是另一個 coupling,所以 πε,cπ,c\langle\pi_\varepsilon,c\rangle\ge\langle\pi^*,c\rangle——entropic 解的運輸成本永遠不低於真實 OT。它多付的成本換來的是 H(πε)H(π)H(\pi_\varepsilon)\ge H(\pi^*):拆得更散。

展開細節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 的偏差

BBμ\mu 的樣本、BBν\nu 的樣本,各成經驗分佈 μ^B,ν^B\hat\mu_B,\hat\nu_B,算 W2(μ^B,ν^B)W_2(\hat\mu_B,\hat\nu_B),對抽樣取期望。兩件事可以說清楚。

距離偏高。 這可以證明:OT(μ,ν)=minππ,c\mathrm{OT}(\mu,\nu)=\min_\pi\langle\pi,c\rangle(μ,ν)(\mu,\nu) 這一對是jointly convex 的(兩個 coupling 的凸組合仍是對應邊際凸組合的 coupling,所以最小值不會高於凸組合),而 E[μ^B]=μ\mathbb E[\hat\mu_B]=\muE[ν^B]=ν\mathbb E[\hat\nu_B]=\nuM5.0 立刻給

E[W22(μ^B,ν^B)]    W22(Eμ^B,Eν^B)=W22(μ,ν).\mathbb E\big[W_2^2(\hat\mu_B,\hat\nu_B)\big]\;\ge\;W_2^2\big(\mathbb E\hat\mu_B,\mathbb E\hat\nu_B\big)=W_2^2(\mu,\nu).

直覺是:組內的配對被「組裡只有這 BB 家店」的限制綁住,每個倉庫的最近選項比全體裡少,被迫送得更遠。極端例子:μ=ν\mu=\nu,真實 W2=0W_2=0;但兩組獨立抽出的 BB 個點幾乎不會重合,W2(μ^B,ν^B)>0W_2(\hat\mu_B,\hat\nu_B)>0 嚴格成立,且在 dd 維以 B1/dB^{-1/d} 的速率緩慢趨近零(經驗分佈的 W2W_2 收斂速率,見 Peyré & Cuturi [2])。所以 minibatch W2W_2 對「兩個分佈其實相同」的判斷有系統性偏差BB 越小、維度越高越嚴重 [4]。

配對偏亂。 把每一組內的最佳配對合起來,得到一個「minibatch coupling」πˉB=E[πB]\bar\pi_B=\mathbb E[\pi^*_B]。它是 μ,ν\mu,\nu 的合法 coupling(每組內邊際都對,平均後也對),所以它的成本 W22\ge W_2^2——又是「任何 coupling 給上界」。而且它不是一個映射:同一個倉庫在不同組裡會被配到不同的店,平均後一個 xx 對應到一片 yyBB\to\inftyπˉBπ\bar\pi_B\to\pi^*B=1B=1πˉ1=μν\bar\pi_1=\mu\otimes\nu(獨立 coupling)——BB 在「精確 OT」與「完全隨機配」之間插值,角色與 Sinkhorn 的 ε\varepsilon 相同。

兩種近似殊途同歸:都給出一個比精確解更分散的 coupling、成本偏高;一個用 ε\varepsilon 控制、一個用 1/B1/B 控制。

回到五萬個倉庫:該用哪個

回答起點問題。第一條路(entropic)成本偏高、配對偏散,偏差由 ε\varepsilon 控制,代價是 O(mn)O(mn) 的矩陣乘向量——五萬對五萬在 GPU 上可行,但 KK2.5×1092.5\times10^9 格的記憶體(可以分塊)。第二條路(minibatch)成本偏高、配對偏亂,偏差由 BB 控制,記憶體只要 B2B^2,但每組的估計有隨機性、要多抽幾組平均。

選擇的判準:你要的是「距離」還是「配對」? 若要的是一個數字(兩個大點雲差多少),Sinkhorn divergence 加中等的 ε\varepsilon 通常最好——偏差可控、無隨機性。若要的是「每個 xx 該配哪個 yy」而且要在訓練迴圈裡每一步都算,minibatch 是唯一現實的選擇——這時你要清楚配對是「偏亂」的,BB 越大越接近真實 OT。

依賴的假設:Sinkhorn 的收斂需要 KK 沒有全零的列或欄(ε\varepsilon 太小時會發生,要用 log 域);minibatch 的偏差分析假設樣本獨立同分佈抽自 μ,ν\mu,\nu。兩者的偏差方向(成本偏高)是嚴格的,但偏差的大小要靠實驗量。

回到五萬個倉庫:縮小成兩千個,把三個數並排

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 真值 10.2510.25;用全部 2000 個樣本精確解 assignment 得 10.7310.73(有限樣本的估計本身就偏高一點)。Sinkhorn 成本隨 ε\varepsilon10.7710.77 單調升到 13.3113.31、entropy 從 11.611.6 升到 15.015.0——偏高、偏散,而且 ε=0.05\varepsilon=0.05 就已經在 K=eC/εK=e^{-C/\varepsilon} 下溢的邊緣,所以程式直接寫成 log\log 域。minibatch 從 B=8B=811.9111.91 降到 B=512B=51210.7710.77從上方逼近 10.7310.73

這裡有一個實作上很會騙人的地方,值得自己踩一次。把 iters 從 1000 降回 300,ε=0.05\varepsilon=0.05 會給出 10.4910.49——比精確解 10.7310.73 還低,直接違反上面那個「entropic 解的成本不低於真實 OT」。原因不是定理錯了,是那個 PP 還不是合法的 coupling:300 輪之後列邊際的 L1L_1 誤差還有 0.0380.038,質量沒有守好,成本自然可以被壓到下界以下。ε\varepsilon 越小收斂越慢,所以看到「比精確解還低」的第一件事是去檢查邊際,不是去懷疑定理。

刻意違反一個假設。YY 換成與 XX 同分佈的另一組樣本(真值 W2=0W_2=0):2000 對 2000 的精確解給 0.020.02,minibatch 在 B=8B=8 時給 1.331.33B=512B=512 時給 0.0670.067——它會把「兩個相同的分佈」判成「有差」,BB 越小越嚴重。再把 ε\varepsilon 壓到 10310^{-3} 且改回普通域的 K=eC/εK=e^{-C/\varepsilon}:矩陣大量下溢成 0、K @ v 出現零、除法變成 inf——這就是為什麼要在 log\log 域做。

先消化一下

想一想

你用 Sinkhorn 算兩個點雲的距離,發現把同一個點雲跟自己比,算出來的成本不是零。這是:

想一想

你要在訓練迴圈裡、每一步都拿到「這批 xx 該配哪些 yy」的最佳配對,資料有百萬筆。現實的選擇是:

想一想

把 Sinkhorn 的 ε\varepsilon 調到很大、或把 minibatch 的 BB 調到 1,兩者的極限分別是:

參考文獻

  1. Cuturi, M. Sinkhorn Distances: Lightspeed Computation of Optimal Transport. NeurIPS 2013.(把 entropic 正則化與 Sinkhorn 迭代帶進大規模 OT。)
  2. Peyré, G., Cuturi, M. Computational Optimal Transport. Foundations and Trends in Machine Learning 11(5–6), 2019.(第 4 章:entropic OT、Sinkhorn 的推導與收斂、log 域實作、ε\varepsilon 的偏差量級;第 8.5 節:經驗分佈的收斂速率。)
  3. Sinkhorn, R., Knopp, P. Concerning Nonnegative Matrices and Doubly Stochastic Matrices. Pacific Journal of Mathematics 21(2), 1967.(交替正規化收斂到唯一的對角縮放。)
  4. Fatras, K., Zine, Y., Flamary, R., Gribonval, R., Courty, N. Learning with Minibatch Wasserstein: Asymptotic and Gradient Properties. AISTATS 2020.(minibatch OT 的偏差、μ=ν\mu=\nu 時不為零、隨 BB 的收斂。)
  5. Genevay, A., Peyré, G., Cuturi, M. Learning Generative Models with Sinkhorn Divergences. AISTATS 2018.(Sinkhorn divergence:扣掉自我運輸成本。)