☰
Wasserstein距离:从KL/JS失效到Sinkhorn与WGAN实战
2026/10/1 12:33:34 网站建设 项目流程

我第一次被 Wasserstein 距离说服,不是因为哪个漂亮的定理,而是因为一次实验里 KL 散度直接返回了inf,训练曲线在几万步之后一动不动。当时我在做两批高维特征的分布对齐,一批特征挤在一个窄窄的区域里,另一批在旁边的另一个位置,两者几乎不重叠。按直觉,这两批分布只是"错开了一段距离",应该很容易被拉近;但 KL 散度告诉我它们相距无穷远,JS 散度倒是给了个有限的数,可惜那个数是个常数——不管我把两批分布挪多远,它都雷打不动地对模型说"两者差一个固定的量",梯度为零。Wasserstein 距离就是在这个场景下把我捞出来的东西:它不只看两个分布有没有重叠,还看重叠不上的部分到底隔了多远。

这篇文章我想把 Wasserstein 距离拆到底:它到底是什么、定义里每一项在干什么、为什么一维可以直接排序求解、和 KL/JS/TV 放在一起比强在哪里、真上手算的时候 Sinkhorn 怎么用、代价矩阵怎么设计、以及我自己在项目里踩过的那几个坑。不管你是做生成模型、域自适应、点云配准,还是单纯在统计课上被这个名词卡住,读完都应该能自己动手把它算一遍。

1. 一次对比实验把我从KL散度推向Wasserstein

1.1 两个完全不重叠的分布,KL和JS给出的答案都不讲理

先把场景说清楚。假设 μ 是全部质量压在 0 处的一个分布,ν 是全部质量压在 θ 处的一个分布,θ 是个正数。这两个分布的距离,用肉眼判断显然是 θ——它们隔了 θ 这么远,θ 越大越远,θ 越小越近。但只要 θ 不为零,也就是两个分布完全不重叠,KL 散度给出的答案就是正无穷。原因很直白:KL 的定义里有一个密度比 p/q,当 q 在 p 有质量的地方等于零时,这个比值直接爆炸。KL(μ‖ν) 和 KL(ν‖μ) 还会给出不同的结果,甚至一个是有限一个是无穷。

JS 散度做了个"折中",把两者和中间分布比一遍再平均,于是结果有界了,最大不超过 log 2。问题在于,只要两个分布不重叠,JS 就死死钉在 log 2 这个常数上,θ 从 0.001 变到 1000,JS 纹丝不动。对优化来说这是灾难:梯度恒为零,模型完全不知道该往哪个方向走。TV 距离(全变差)也好不到哪去,它直接就是 1,同样是个"要么重叠要么不重叠"的二值化判断,丢掉了所有的几何信息。

这三个度量的共同毛病是:它们本质上在数"重叠面积",而不是在量"隔了多远"。而 Wasserstein 距离恰恰是从"隔了多远"这个角度定义的,这也是它最核心的价值。

1.2 把分布看成沙堆,距离就变成了搬运的力气

Wasserstein 距离最经典的解释叫"推土机距离"(Earth Mover's Distance)。想象你面前有两堆沙子,一堆是分布 μ 的形状,另一堆是分布 ν 的形状。你的任务是只用一把铁锹,把第一堆沙子一点点搬到第二堆的形状上去,最后让两堆沙子的轮廓完全一致。

搬运是有代价的:搬一单位的沙子,走 d 的距离,就消耗 d 的力气(如果按 1 次方计)。那么最优策略下的总消耗,就是这两个分布之间的 W₁ 距离。这个类比妙就妙在,它同时抓住了三件事:第一个是"有多少沙子要搬",对应两个分布质量上的差异;第二个是"每一铲要搬多远",对应两个分布在地理上的错位;第三个是"怎么搬最省",这是整个最优传输问题的核心,也是 Wasserstein 距离求解的难点所在。

回到刚才的例子:μ 是 0 处的一整堆沙子,ν 是 θ 处的另一堆。搬法只有一种,把所有沙子整体平移 θ,消耗就是 θ。所以 W₁(μ, ν) = θ。这个答案既有界,又连续,还随 θ 线性变化——梯度稳定地告诉你"往左挪"。这就是当年 WGAN 用它替换 JS 散度的根本动机。

1.3 这份"搬运成本"能用在哪些地方

我自己实际接触过的场景有这么几类。生成模型里,WGAN 系列用 Wasserstein 距离当作训练目标,绕开 JS 散度的梯度消失问题;域自适应里,用它衡量源域和目标域特征分布之间的差异,作为一个可微的正则项;点云处理里,用它做形状匹配和配准,因为点云天然就是离散测度;单细胞数据分析里,用它做不同时间点样本之间的分布对齐,进而推断细胞状态的变化轨迹;图像处理里的颜色迁移,本质就是求两个颜色分布之间的传输方案。此外还有分布鲁棒优化、公平性度量、文本嵌入空间的分布比较等等。

这些应用有一个共同点:它们关心的都不是"两个分布像不像",而是"两个分布之间要付出多大代价才能互相转化"。如果你的任务里出现了这个味道,Wasserstein 距离大概率就是你要找的工具。

2. 把"搬沙子"写成数学:传输计划与Wasserstein的完整定义

2.1 传输计划就是一张调度表

光有直觉不够,得把"搬运方案"写成数学对象。我们用一个联合分布 π(x, y) 来表示方案:π(x, y) 的数值代表"从 x 处搬多少质量到 y 处"。这个 π 必须满足两个约束,也就是它的两个边缘分布要分别等于 μ 和 ν:

  • 把所有 y 方向的 π(x, y) 积分掉,得到的是 μ(x),意思是"从 x 出发的总质量,必须等于 μ 在 x 处的质量",沙子不能凭空产生。
  • 把所有 x 方向的 π(x, y) 积分掉,得到的是 ν(y),意思是"到达 y 的总质量,必须等于 ν 在 y 处需要的质量",沙子也不能凭空消失。

把所有满足这两个约束的 π 收集起来,构成一个集合,记作 Π(μ, ν),叫传输多面体。对其中任何一个 π,我们都能算出它的总运输成本——把每一小块质量乘上它要走的距离,再全部加起来。最优传输问题就是在这个集合里找一个让总成本最小的 π。

有了这个框架,Wasserstein 距离的定义就顺理成章了:

W_p(μ, ν) = ( inf over π ∈ Π(μ, ν) of ∫ d(x, y)^p dπ(x, y) )^(1/p)

拆开看:d(x, y)^p 是单位质量从 x 搬到 y 的代价(距离的 p 次方),dπ(x, y) 是这条路线上的质量,积起来就是总代价;inf 表示在所有可行方案里挑最省的那个;最后开 p 次方根,把量纲拉回到"距离"本身。整个定义没有一个多余的符号。

2.2 手算一个2×2的例子,把公式彻底拆开

抽象的定义容易让人飘,我们算一个最小的例子。设 μ 的质量是:0 处 0.7,2 处 0.3;ν 的质量是:1 处 0.4,3 处 0.6。用 W₁(也就是 p=1)来算,距离用绝对差。

先写代价矩阵,行是 μ 的支撑点,列是 ν 的支撑点:

从 \ 到y=1y=3
x=013
x=211

传输计划有四个未知量 π₁₁、π₁₂、π₂₁、π₂₂,约束是行和等于 μ、列和等于 ν。四个未知量加四个约束,但约束里有一组是冗余的(行和加起来等于列和加起来等于总质量 1),所以实际只有三个独立的自由度,解是一维的,可以参数化成一个变量 t = π₁₁。

代入约束:π₁₂ = 0.7 − t,π₂₁ = 0.4 − t,π₂₂ = t − 0.1。非负要求把 t 限制在 [0.1, 0.4] 区间内。总代价写成 t 的函数:

代价 = t·1 + (0.7 − t)·3 + (0.4 − t)·1 + (t − 0.1)·1 = 2.4 − 2t

代价随 t 递减,取 t = 0.4 时最小,最小代价是 1.6。此时的最优方案是 π₁₁ = 0.4、π₁₂ = 0.3、π₂₁ = 0、π₂₂ = 0.3。所以 W₁ = 1.6。

这个结果有个值得记住的特征:最优方案里 π₂₁ = 0,也就是说没有沙子从 x=2 搬到 y=1。整个方案是"有序"的——排序后第一个点的质量送到排序后第一个点,不会出现交叉搬运。这不是巧合,而是一维情形下的必然结论,下一节会展开。

2.3 p次方与开根号:定义里那一层到底在干什么

很多人第一次看定义会疑惑,为什么非要套一个 p 次方再开根号,直接定义成最小总代价不就行了。原因是度量公理。要求 p ≥ 1,开完根号之后的 W_p 才满足三角不等式,也就是 W(μ, ω) ≤ W(μ, ν) + W(ν, ω)。如果取 p < 1,三角不等式会被破坏,它在严格意义上就不叫"距离"了。

p 具体取几,取决于你想强调什么。p = 1最符合直觉,量纲就是普通距离,对异常值也更稳健,因为代价是线性增长的而不是平方增长。p = 2的代价对远距离的错配惩罚更重,它会倾向于把大的偏移摊平,得到的传输方案在几何上更"平滑",因此被广泛用于梯度流、Wasserstein 重心和形状分析。工程上如果只是想要一个好的差异度量,p = 1 和 p = 2 都用,看你的代价函数是否希望放大远距离误差。

3. 连续分布、耦合视角与一维闭式解

3.1 耦合:把两个分布"绑"成一个联合分布

上一节的 π 还有另一个身份:它是一个耦合(coupling)。换一种说法,取一对随机变量 (X, Y),让 X 服从 μ、Y 服从 ν,那么 (X, Y) 的联合分布就是一个耦合。在这个视角下,Wasserstein 距离变成:

W_p(μ, ν) = ( inf over joint distributions of (X, Y) with marginals μ, ν of E[ d(X, Y)^p ] )^(1/p)

这个写法比积分的写法更容易理解。它在说:我要在两个分布的所有"配对方式"里,选一种让 X 和 Y 的平均距离最小的配对方式。耦合里可以存在随机性——X 可以按概率分给不同的 Y;但如果问题足够好(一维、代价函数是凸的),最优耦合会是确定性的,也就是一个从 X 到 Y 的函数映射,这正是 18 世纪 Monge 最初提出的形式;Kantorovich 后来把它松弛成线性规划,才有了可以求解的版本。这个松弛是整个领域的转折点。

理解耦合对做工程很重要,因为你实际拿到的"距离值"取决于你允许哪些配对。如果你偷偷加了结构约束,比如只允许相邻的点互相搬运,那你算出来的已经不是标准的 Wasserstein 距离了,而是某个受限版本。这个区别在做算法选型时经常被忽略。

3.2 一维为什么可以直接排序求解

一维是 Wasserstein 距离的"天堂"。结论是:只要分布在一维直线上,最优配对一定是保序的——把 μ 的样本从小到大排,ν 的样本从小到大排,第 k 个配第 k 个,不需要解任何优化问题。用分位数函数(累积分布函数的反函数)写出来更漂亮:

W_p(μ, ν)^p = ∫₀¹ | F_μ⁻¹(u) − F_ν⁻¹(u) |^p du

当 p = 1 时,这个式子还可以进一步在离散情况下化成"两个累积分布函数曲线之间的面积"。原因是这样的:如果存在两条交叉的搬运路线,也就是 x₁ < x₂ 但配对了 y₁ > y₂,把它们互换一下,总路程一定不会变长。反复做这种"消交叉"操作,最后一定得到完全有序的配对。这个论证思想简单但结论很强,它意味着一维的 Wasserstein 距离可以在 O(n log n) 时间内精确算出来——比高维的线性规划快了不止一个数量级。

回到 2.2 节的手算例子验证一下。μ 的分位数函数:u < 0.7 时取值 0,u ≥ 0.7 时取值 2。ν 的分位数函数:u < 0.4 时取值 1,u ≥ 0.4 时取值 3。分段积分:0 到 0.4 段差值是 1,宽度 0.4;0.4 到 0.7 段差值是 3,宽度 0.3;0.7 到 1 段差值是 1,宽度 0.3。加起来 0.4 + 0.9 + 0.3 = 1.6,和线性规划的结果完全吻合。

3.3 维数一高就天翻地覆:样本复杂度

一维的顺利会让人产生一种错觉,以为高维也差不多。实际差距非常大。Wasserstein 距离的估计精度衰减速度高度依赖维度:在一维,用 n 个样本估计的经验 Wasserstein 距离,误差以接近 n 的倒数速度收敛;到了三维以上,误差量级大约是 n 的 −1/d 次方。也就是说,维度每增加一点,你就需要成倍成倍地加样本才能维持同样的精度。这个现象通常被称为最优传输的维数灾难。

维度 d达到相同精度所需样本量级直观感受
1n几乎不心疼
2约 n^(1/2) 的相对量级还能接受
4约 n^(1/4) 的相对量级开始难受
10约 n^(1/10) 的相对量级基本不可用

这张表解释了一个常见的困惑:为什么论文里 Wasserstein 距离效果很好,我在自己的任务上却感觉不到。很可能是因为你的特征维度偏高,而样本量远远不够。这时候通常需要考虑施加更强的结构假设(比如切片、低秩、或者用 Sinkhorn 散度替换),后面第 6 节会讲。

3.4 一维的两种实现,附上对照说明

第一种是标准的分位数积分法,用两个累积分布的断点把 [0, 1] 区间切细,逐段累加。

import numpy as np def wasserstein_1d(x, y, p=1.0): """一维经验分布的精确 W_p。x, y 是两组样本,允许长度不同。""" x = np.sort(np.asarray(x, dtype=float)) y = np.sort(np.asarray(y, dtype=float)) n, m = x.size, y.size # 所有分位数函数的跳变点 cuts = np.unique(np.concatenate([np.arange(n + 1) / n, np.arange(m + 1) / m])) total = 0.0 for lo, hi in zip(cuts[:-1], cuts[1:]): mid = 0.5 * (lo + hi) # F^{-1}(u) = inf{t : F(t) >= u},对应下标 ceil(u*n)-1 ix = min(int(np.ceil(mid * n - 1e-9)) - 1, n - 1) iy = min(int(np.ceil(mid * m - 1e-9)) - 1, m - 1) total += abs(x[max(ix, 0)] - y[max(iy, 0)]) ** p * (hi - lo) return total ** (1.0 / p)

第二种是通用的离散线性规划,可以处理非均匀权重,还能顺便把最优传输计划取出来看。

import numpy as np from scipy.optimize import linprog def wasserstein_lp(x, a, y, b, p=1.0): """返回 (W_p, 最优传输计划)。x/y 为支撑点,a/b 为对应质量。""" x, y = np.asarray(x, float), np.asarray(y, float) a, b = np.asarray(a, float), np.asarray(b, float) n, m = x.size, y.size C = np.abs(x[:, None] - y[None, :]) ** p # 代价矩阵 A_eq, rhs = [], [] for i in range(n): # 行和等于 a row = np.zeros(n * m); row[i * m:(i + 1) * m] = 1.0 A_eq.append(row); rhs.append(a[i]) for j in range(m): # 列和等于 b row = np.zeros(n * m); row[j::m] = 1.0 A_eq.append(row); rhs.append(b[j]) res = linprog(C.ravel(), A_eq=np.array(A_eq), b_eq=np.array(rhs), bounds=(0, None), method="highs") plan = res.x.reshape(n, m) return res.fun ** (1.0 / p), np.round(plan, 6) print(wasserstein_lp([0, 2], [0.7, 0.3], [1, 3], [0.4, 0.6], p=1.0)) # (1.6, array([[0.4, 0.3], # [0. , 0.3]]))

两段代码跑出来的结果和手算一致。第一种适合数据量大的一维场景,第二种适合小规模、需要完整方案的高维场景,代价是变量数量是支撑点数量的乘积,点一多内存就撑不住了。

4. 把KL、JS、TV和Wasserstein摆到同一张表里

4.1 不重叠场景下的数值对比

还拿 μ = 在 0 处的点质量、ν = 在 θ 处的点质量做例子,θ > 0。把四个度量放在一起:

度量对称性是否满足三角不等式不重叠时的取值随 θ 变化
KL(μ‖ν)否否+∞无意义
JS是是(其平方根)log 2(常数)不变
TV是是1(常数)不变
W₁是是θ线性增长
W₂是是θ线性增长

这张表最该被记住的一行是最后两行。当两个分布不重叠时,前三个度量集体失效——要么无穷,要么常数,梯度信号完全丢失。而 Wasserstein 距离仍然给出一个随几何位置线性变化的有限值,这个值直接就是"还差多远"。这就是它在生成模型里不可替代的原因。

4.2 对偶形式与Lipschitz约束的由来

W₁ 有一个特别实用的对偶形式,叫 Kantorovich-Rubinstein 对偶:

W₁(μ, ν) = sup over f with Lip(f) ≤ 1 of [ E_μ f − E_ν f ]

意思是,Wasserstein 距离等于"在所有 1-Lipschitz 函数里,找一个能让两个分布的期望差值最大的那个函数,然后取这个最大差值"。这个式子的意义很大:它把一个求最小值的线性规划,变成一个求最大值的优化问题,而后者可以直接用神经网络参数化来近似。WGAN 里那个被称为 critic 的网络,干的就是这件事——它去逼近那个最优的 f,而 Lipschitz 约束就是这条对偶定理里的硬性要求。

我第一次读到这里时最困惑的就是:为什么非得是 1-Lipschitz?直观解释是这样的:如果对 f 不加约束,你只要让 f 在 μ 支撑集上取很大的正值、在 ν 支撑集上取很大的负值,差值可以无限大,这个上确界就没有意义了。加上 Lipschitz 约束之后,f 的变化速度被限住,它就只能"老老实实"地反映分布之间的几何差异。约束一旦放松,整个估计就失去意义——这也是我在 WGAN 项目里栽过的坑,后面细说。

4.3 测地线与重心:两个分布之间的"中间态"

Wasserstein 空间还有一个 KL 散度完全没有的性质:它是个长度空间,两点之间有"直线"。取最优传输计划 π,定义中间分布 μ_t 为"把 π 里每条搬运路线都只走 t 的比例"所得到的分布。t 从 0 走到 1,你会看到沙堆从 μ 的形状平滑地变形到 ν 的形状,中间那些形状就是测地线上的点。

这个性质在应用里非常直接。图像插值可以用它做:把两张图像的颜色分布或像素分布当成两个测度,沿测地线取中间态,得到的过渡帧比直接线性混合自然得多,因为线性混合会让物体出现"半透明重影",而最优传输是整体搬运,不会产生鬼影。Wasserstein 重心(把多个分布平均)也是同一套逻辑的推广,在颜色迁移、形状平均、多域数据混合里都有落地。

5. 真要算的时候:从线性规划到Sinkhorn迭代

5.1 精确解的内存与时间墙

拿通用线性规划直接解最优传输,变量数量是 n × m。如果两组数据都是 10000 个点,代价矩阵就有 10⁸ 个元素,用 float64 存需要 800 MB,还没开始算内存就炸了。时间上,经典内点法或网络单纯形法的复杂度大致在 O(n³ log n) 量级,n 上万时基本没有可行性。

所以实际工程里的第一道选择题是:你到底需不需要精确解。如果支撑点只有几百个(比如图像的 256 级颜色直方图,或者小规模的点云关键点),精确 LP 完全够用,简单可靠。如果点的数量上万甚至更多,就必须转向近似方法。

5.2 熵正则化:把线性规划摊成矩阵缩放

最主流的近似方法是熵正则化。它的思路是在原来的最小化目标后面加一项负熵:

min over π of [ <C, π> − ε H(π) ]

负熵这一项会让 π 变得平滑、分散,不再死盯着少数几条最优路线。这样做有两个好处:一是把一个可能不唯一的线性规划变成了严格凸问题,解唯一且稳定;二是解的形式特别规整,可以写成两个对角矩阵夹住一个核矩阵的形式:

π = diag(u) · K · diag(v),其中 K = exp(−C / ε)

于是求解就变成交替更新 u 和 v,也就是所谓的Sinkhorn 迭代:先用当前的 v 更新 u,再用新的 u 更新 v,反复几轮,行和列和就会收敛到 μ 和 ν。整个过程就是几次矩阵乘法,向量化之后在 GPU 上跑得非常快。

import numpy as np def sinkhorn(a, b, C, eps=0.05, iters=3000, tol=1e-10): """熵正则化最优传输,返回传输计划和正则化后的总代价。""" K = np.exp(-C / eps) u = np.ones_like(a) v = np.ones_like(b) for _ in range(iters): v_new = b / (K.T @ u + 1e-300) u_new = a / (K @ v_new + 1e-300) if np.max(np.abs(u_new - u)) < tol: u, v = u_new, v_new break u, v = u_new, v_new pi = u[:, None] * K * v[None, :] return pi, np.sum(pi * C)

Sinkhorn 迭代的收敛速度是线性的,收敛率取决于 ε 和代价矩阵的尺度——ε 越小,收敛越慢,需要的迭代次数越多。实践中常见的做法是设定一个固定迭代次数(几百到几千),而不是等它收敛,因为下游任务往往只需要一个足够好的近似。

5.3 epsilon 怎么选,数值下溢怎么救

ε 是这套方法里唯一需要认真调的参数,我把它对应的行为整理成一张表:

ε 的取值解的性质收敛速度数值风险
偏大(如代价量级的 0.1 倍)非常平滑,偏离真实 W 较多快,几十轮即可低
适中(如代价量级的 0.01 倍)较好的折中中等中
偏小(如代价量级的 1e-4 倍)逼近真实 W很慢,需上万轮高,K 下溢为 0

当 ε 很小的时候,exp(-C/eps)里的指数会变成很大的负数,K 直接下溢成 0,除以零就得到 NaN。救法有两条:一是把整个迭代搬到对数域,用 logsumexp 的写法替代直接乘除,把数值稳定的责任交给对数运算;二是做一个预处理,把代价矩阵整体缩放到 [0, 1] 或者除以它的中位数,让 ε 的相对量级有意义。我自己的经验是,永远先归一化代价矩阵再选 ε,否则同一份代码换一个数据集就崩。

另外还有一个坑:熵正则化会引入系统性偏差——即使输入的是同一个分布,OT_ε(μ, μ) 也不等于零,因为负熵项在鼓励质量分散。修正办法是用 Sinkhorn 散度:把 OT_ε(μ, ν) 减去 OT_ε(μ, μ) 和 OT_ε(ν, ν) 各自的一半。这样得到的量在 μ = ν 时严格为零,而且自对偶性更好。如果你的下游任务对"距离为零"这个语义敏感,务必做这步修正。

6. 我踩过的坑和几个真实落地的用法

6.1 WGAN里的符号方向与Lipschitz约束,两个静默失败点

第一次实现 WGAN 的时候,我照着对偶公式写了 critic 的损失,训练几轮后发现生成器完全没有改善,但两条 loss 曲线都在下降,看不出任何异常。问题出在符号方向:critic 是在最大化E_ν f − E_μ f,生成器是在最小化同一个量,两个优化目标方向相反。如果实现时把 critic 的符号写反了,critic 会变成在最小化它本该最大化的量,loss 看起来正常,实际学的是一个无意义的函数。排查方法很土但有效:固定一批数据,手动算一次 E_ν f − E_μ f 的值,看它的正负和量级是否符合预期。

第二个坑是 Lipschitz 约束的实现方式。早期用权重裁剪,把网络参数硬截断到一个小范围,结果参数大量堆积在边界上,critic 的表达能力被削掉一大半。后来改用梯度惩罚,把 f 的梯度范数往 1 上拉,稳定性明显好得多。这里的关键认识是:Lipschitz 约束不是正则化小花招,它是那条对偶等式成立的前提。约束一旦失效,你算出来的就不再是 Wasserstein 距离,而是一个没有下限的奇怪数值。

6.2 代价矩阵才是业务知识落地的地方

我见过太多人把代价矩阵默认设成欧氏距离就开始跑,然后抱怨效果不好。代价矩阵 C(x, y) 是"从 x 搬到 y 要花多少代价"的建模,它是整条链路里唯一能注入先验知识的地方,值得反复雕琢。

举两个我做过的例子。点云配准里,除了空间欧氏距离,我还把法向量夹角加了进去,代价写成空间距离加上一个权重乘以法向夹角的正弦值,这样配准结果不会把两个朝向相反的局部平面强行对上。颜色迁移里,直接在 RGB 空间算距离效果一般,因为 RGB 的欧氏距离和人的感知不一致,换到 Lab 空间之后,同一个代价矩阵的效果立刻好很多。还有一点必须强调:先做坐标归一化。如果某一维的数值范围是另一维的一万倍,代价矩阵会被这一维完全主导,其他维度等于被忽略了。我一般会先把每维特征标准化到单位尺度,再调权重。

6.3 小批量OT不等于真OT,切片是一个务实的退路

做生成模型时特别容易犯的一个错,是把一个 batch 内的样本当成整个分布,直接算 Wasserstein 距离当损失。这样算出来的量叫做小批量最优传输,它和真实分布之间的 Wasserstein 距离之间有一个不易察觉的偏差。原因很直观:batch 里的样本都是有限的,最优传输可以在这些有限的点上"钻空子",找出一条比真实最优方案更便宜的路线,因为真实分布中本该参与搬运的很多点根本不在这个 batch 里。结果就是,batch 越小,这个低估越严重。

再加上第 3.3 节提到的维数灾难,高维场景下小批量估计的方差会大得离谱,用两个 batch 算出的距离去指导梯度更新,噪声可能比信号还大。这时候我会考虑切片 Wasserstein:把高维分布沿很多随机方向投影到一维,在每个方向上用分位数公式精确求解,最后取平均。它的好处是绕开了维数灾难,收敛速度回到一维的量级,而且实现只有几十行。代价是它只沿方向比较边缘分布,会丢失一些联合结构的差别,所以只适合作为起点或者辅助项,不能指望它刻画全部几何。

6.4 质量不守恒时,不平衡和部分传输怎么想

标准的最优传输有一个硬前提:两边的总质量必须完全相等,而且要一比一地搬完。真实数据经常不满足。比如比较两种实验条件下的单细胞样本,两边的细胞数差了好几倍;再比如做图像域自适应,目标域里有一大块区域在源域里根本找不到对应物。强行套标准公式,结果会被那部分多余的质量严重扭曲,因为它必须被硬塞给某个目标点,代价极大。

处理方法有两类。不平衡最优传输是在约束上放松,允许边缘分布有一定偏差,用 KL 散度对边缘偏差加惩罚,惩罚系数控制放松程度。部分传输则是允许一部分质量直接销毁或从无到有生成,代价设为一个常数阈值。我选择的时候看一个判断:多余的质量是"噪声"还是"有意义的新模式"。如果是噪声(比如离群点),用部分传输,让它被丢掉;如果是真实的新模式(比如目标域特有的一个类别),用不平衡传输,让边缘上的偏差被保留下来。

判断参数调得对不对,有个很实用的小检查:跑完之后统计一下有多少质量最终没有被匹配上,也就是传输计划的行和列和与原始质量之间的差。如果这个比例大得超出预期,说明惩罚给得太松;如果几乎为零,而你明明知道数据里有离群点,说明惩罚太紧,那部分离群点正在污染整个传输方案。这个检查比盯着最终指标有用得多,因为它直接告诉你模型在数据上做了什么,而不是只给你一个数。

最后分享一个我自己的习惯:任何一次用 Wasserstein 距离的实验,我都会先把支撑点数压到几百以内,用精确的线性规划跑一遍,把最优传输计划打印出来看一眼。看看哪些点搬到了哪些点、有没有出现明显违反常识的配对。这一步花不了十分钟,但它能帮我在调 Sinkhorn 的 ε 和代价矩阵权重之前,先确认问题建模本身是对的。很多所谓"算法不work"的情况,回头看都是代价矩阵写错了。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询