ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Wasserstein距离与最优传输:从搬土直觉到Sinkhorn实战

Wasserstein距离与最优传输:从搬土直觉到Sinkhorn实战 1. 为什么常规的分布度量会失灵Wasserstein距离这几年在各种论文、开源库、技术分享里出现得越来越频繁尤其是做生成模型、域适应、图像检索的朋友几乎绕不开它。但真正把它讲明白、算清楚的人不多多数资料一上来就甩出最优传输和耦合这些词把刚接触的人劝退。我一开始也是被它绕得快放弃后来在一个分布对齐的项目里被迫从头啃才算摸清了它的直觉内核。这篇就把我当时踩过的路、想通的点和真正落地的代码摊开来聊。先说清楚它解决什么问题当你手里有两个概率分布怎么衡量它们差多远这个问题看着简单实际在工程里非常要命。比如你在做用户行为建模训练集的分布和线上真实分布不一致再比如做图像风格迁移源域和目标域的特征分布对不上又或者你在训练生成网络需要告诉模型你生成的分布离真实分布还有多远。这些场景都需要一个数值能稳定、准确地反映两个分布的差异。Wasserstein距离就是为此而生的度量它比KL散度、JS散度更讲道理尤其在两个分布几乎不重叠的时候依然能给出有意义、有梯度的信号。适合谁看如果你做过GAN、做过域适应、做过任何跟分布比较相关的工作或者单纯被这个名词勾起了好奇这篇都值得读完。我尽量不用拗口的数学符号堆砌而是从搬土这件事讲起把为什么、怎么算、怎么用一层层拆开。读完之后你至少能做到三件事第一向同事用大白话解释清楚它的直觉第二手算一维情况下的结果第三用Python把它、算出来并理解每种实现的取舍。1.1 从两张直方图说起假设我拿两组用户年龄数据A组分布集中在25到35岁B组集中在40到50岁两组几乎不重叠。我想用一个数字量化它们有多不一样。最朴素的做法是逐桶比较高度加个绝对值求和这就是所谓的L1距离或者全变差距离的简化版。它有用但有个明显毛病它只关心同一个年龄桶里的人数差不关心隔壁桶的人是相近的。A组集中在30岁、B组集中在32岁和A组集中在30岁、B组集中在48岁在某些度量下算出来的差异可能差不多但直觉上前者应该小得多。人是有连续性的30岁和32岁接近30岁和48岁差得远好的度量应该把这种距离感纳入进去。Wasserstein距离的出发点就在这里——它衡量的是把一个分布搬成另一个分布所需的最小代价而搬运距离越远代价越高天然就带上了对位置远近的感知。1.2 KL散度与JS散度的结构性软肋做机器学习的人最熟的是KL散度。它的定义是把p分布对q分布的比值取对数期望衡量的是用q去编码p要多付出的信息量。问题在于它不对称而且对零覆盖极其敏感只要q在某个位置上概率为0而p在那边有概率KL就炸成无穷大。这在实际数据里经常发生因为真实分布往往有长尾模型分布覆盖不到一算KL就无穷梯度直接消失。JS散度是KL的对称化版本把两个分布混合起来各算一次KL取平均它有界、对称看起来温和很多但代价是当两个分布完全不重叠时JS散度是常数梯度为零。这一点在GAN的训练里是致命的判别器一旦太强把真实分布和生成分布完全分开JS散度就不再提供任何有效梯度生成器学不动。这正是WGAN出现的重要动机之一而WGAN的核心恰恰就是换成了Wasserstein距离。1.3 一个好的分布度量该具备哪些素质我把当年梳理的判断标准列一下你也可以拿去评估其他度量。第一是几何敏感性距离应该随分布位置的移动平滑变化而不是只看重叠区域。第二是梯度友好尤其在优化场景里即使两个分布不重叠也应该提供有方向、有大小、不消失的梯度。第三是对称性衡量差异不应该因谁是参照而截然不同虽然Wasserstein本身是对称的但它的对偶形式可以让你按需构造出非对称的变体。第四是可计算性理论上再漂亮如果算不动就没有工程价值这也是后面要重点讲Sinkhorn这类近似方法的原因。对照下来KL散度满足部分却栽在梯度上JS散度对称却有平台只有Wasserstein距离在这几条里表现最均衡代价是它的定义背后是一个优化问题计算不是闭式那么简单。2. 推土机视角Wasserstein距离的直觉内核理解Wasserstein距离最好的比喻就是搬土。这个比喻在业内流传很广确实好用但很多人只停在表面没往深里想。我把它掰开讲保证你听完能自己复现这个逻辑。想象有两个土堆形状分别对应两个概率分布。你想把第一个土堆改造成第二个土堆的形状。你要做的是把土从一些位置铲起来运到别的位置放下。每一次搬运都有代价代价等于搬运的土量乘以搬运的距离。Wasserstein距离就是在所有可能的搬运方案中找到总代价最小的那一个这个最小总代价就是两个分布之间的Wasserstein距离。这个定义的美妙之处在于它把分布差异转化成搬土的最省力方案直觉上完全自洽。土堆靠得近稍微推一推就成型代价小土堆隔得远就得大动干戈代价大。位置的远近直接体现在代价里这就是前面说的几何敏感性。更关键的是这个比喻天然解释了两个不重叠分布的情况。土堆A在左边、土堆B在右边完全不重叠KL散度说无穷大JS散度说常数、没梯度而Wasserstein距离会说你至少得把这些土从左边搬到右边距离是D所以代价是M乘以D。它给出了一个明确的、随距离变化的数值梯度也指向把土往右挪这个正确方向。这就是为什么WGAN能用它训练得又稳又不塌。2.1 离散情形下的搬运方案怎么理解把比喻落到数字上。假设一维直线上分布p在位置1有一份土、在位置2有一份土分布q在位置3有一份土、在位置4有一份土。搬土方案有很多种你可以把p在位置1的土搬到3、位置2的土搬到4总代价是1加1等于2平均每份代价1也可以把位置1的搬到4、位置2的搬到3代价是3加1等于4。显然前者更省。还有别的组合方式我们要找的就是最省的那个。这里因为两边各两份、位置错开最优方案并不唯一但最小值是确定的。Wasserstein距离取的就是这个全局最小值。注意方案本身是一个联合分布或者叫传输矩阵它的每个元素表示从源点i搬多少土到目标点j行和等于源分布列和等于目标分布这就是所谓的传输方案必须满足的边际约束。理解了这一层你就懂了为什么它的数学形式里会出现联合分布这个对象。2.2 一维情况下的手算实例一维有个特别漂亮的性质最优传输方案一定跟排序有关。直观理解是左边的土应该尽量往左边的坑搬不该交叉跨越否则交换一下两个目标点就能减少总距离。严格地说一维下把两个分布各自按位置排序然后一一配对搬运就是最优的。我们拿个具体例子手算一下。设分布p在位置0处有0.4的概率质量、位置1处有0.6分布q在位置2处有0.4、位置3处有0.6。一维排序配对0.4份从0搬到2距离2贡献0.4乘2等于0.80.6份从1搬到3距离2贡献0.6乘2等于1.2。总距离等于2.0。这就是W1距离。再换个例子p在0和2各0.5质量q在1和3各0.5质量。排序配对0的0.5搬到1距离1贡献0.52的0.5搬到3距离1贡献0.5总距离1.0。你能明显感觉到分布整体平移了1个单位W1距离就是1完全符合直觉。一维还有个等价写法W1距离等于两个累积分布函数CDF差值的绝对值的积分。这个结论非常实用因为它把优化问题变成了一次积分算起来快得飞起。我在做时序数据分布比较的时候经常直接用这个CDF积分公式几行代码搞定比任何通用求解器都快。2.3 它对位置的感知为何如此关键把上面的例子和JS散度对比差距一目了然。p在0、2各0.5q在1、3各0.5两者完全不重叠。JS散度在这里是常数不管q整体平移到哪里只要不重叠JS散度都一样模型得不到该往哪个方向挪的信息。而Wasserstein距离会随平移量线性增长平移1个单位是1平移2个单位就是2梯度指向减小平移的方向。这种距离即平移量的性质让它在训练过程中提供了极其稳定的信号。我做WGAN的时候最直观的感受是损失曲线平滑得多不再像传统GAN那样动不动震荡爆炸原因就在于即便判别器分得很开损失依然提供有用梯度不会饱和。这不是玄学就是度量本身的数学性质决定的。3. 数学形式从最优传输到标准定义上一节讲的是直觉这一节把直觉翻译成数学。别怕我会把每个符号都说清楚你读完再回头看公式就不会发懵。先约定设Ω是一个度量空间比如实数轴或者欧氏空间两点之间用度量d(x, y)表示距离。给定两个概率分布μ和ν定义在Ω上。我们想找的是一个联合分布π它的第一个边缘是μ第二个边缘是ν。这个π就是传输方案或者叫耦合含义是从x处搬走多少、到y处放多少的全部安排。π(x, y)这个联合分布本身就编码了搬运计划。然后定义代价对每种传输方案π计算搬运的总代价也就是在所有(x, y)上对距离求期望积分∫ d(x, y) dπ(x, y)。Wasserstein距离就是所有满足边缘约束的π里这个代价的最小值。写成公式就是W(μ, ν) inf { ∫ d(x, y) dπ(x, y) : π 的边缘是 μ 和 ν }。这里的inf表示下确界通俗讲就是最小的那个。这一条式子就是Wasserstein距离的全部核心剩下的推广都是在这个骨架上加参数。3.1 传输方案与耦合的严格含义再深挖一层边缘是μ和ν这句话。联合分布π定义在乘积空间Ω乘Ω上它的第一个边缘分布意思是把π对所有y积分掉得到的关于x的分布必须等于μ第二个边缘意思是把π对所有x积分掉得到的关于y的分布要等于ν。这两个约束就是搬运必须恰好把源土堆清空、恰好把目标坑填满的数学表述。这也解释了为什么传输方案不是随便一个联合分布而是边缘固定的那一类。至于为什么最优解一定存在在紧集和连续代价函数下这是有保证的工程上几乎所有情况都满足不用纠结。另外如果μ和ν都是离散的、各n个点那么π就是一个n乘n的矩阵行和是μ的各点质量列和是ν的各点质量这就是大名鼎鼎的传输矩阵。求Wasserstein距离此时变成一个线性规划问题后面代码部分会具体实现。3.2 p阶Wasserstein距离的推广刚才的定义里代价用的是距离d的一次方。推广一下把代价换成距离的p次方再开p次根就得到p阶Wasserstein距离记作W_p。公式是W_p(μ, ν) ( inf { ∫ d(x, y)^p dπ } )^(1/p)。为什么要开根号是为了让量纲和距离一致保证它满足三角不等式成为一个真正的度量。常用的有W1和W2。W1对应p等于1代价是距离本身就是一维CDF积分那个漂亮的形式。W2对应p等于2代价是距离平方它和欧氏空间的几何、高斯分布之间的闭式解联系紧密。这里有个特别实用的结论两个均值分别为m1、m2、协方差分别为Σ1、Σ2的高斯分布它们之间的W2距离有闭式表达等于两者均值差的平方加上协方差矩阵差的一个迹项。做高斯过程、做特征分布对齐的时候这个闭式解经常被拿来直接用因为不需要任何迭代。我早期做域适应的一个baseline就是用高斯假设加上W2闭式解简单但效果不差。3.3 Kantorovich对偶与KR距离原始的Wasserstein定义是个最小化问题带约束算起来不轻松。数学上有个对偶理论能把在联合分布上求最小转化成在函数上求最大。具体来说W1距离等于对所有满足Lipschitz常数不超过1的实值函数f求 E_μ[f] 减 E_ν[f] 的上确界。这就是著名的Kantorovich-Rubinstein对偶记作 KR 距离。这个形式为什么重要因为它把对分布的操作变成了对函数的操作而函数可以用神经网络参数化。这正是WGAN的理论基础用神经网络去逼近那个满足Lipschitz约束的函数叫critic最大化它的期望差就得到了W1距离的估计。为了让神经网络满足Lipschitz约束工程上有几种做法权重裁剪把参数限制在一个小范围、梯度惩罚惩罚梯度范数偏离1、谱归一化约束每一层的谱范数。我实测下来梯度惩罚和谱归一化比权重裁剪稳得多裁剪那个方法对超参太敏感稍微一调就崩。对偶形式还有个好处是它天然给出了往哪个方向挪的梯度信号这就是前面反复强调的优化友好性的数学根源。4. 动手实现三条计算路径与代码理论讲够了来点能跑的。Wasserstein距离的工程实现我总结成三条路一维走排序或CDF积分小规模离散走线性规划大规模连续走Sinkhorn近似。每条路的适用场景、精度、速度差别很大选错了要么慢得跑不动要么精度不够。4.1 一维排序法与CDF积分一维是最省事的直接用累积分布函数差值的积分。离散情况下把两组样本各自排序对应元素相减取绝对值的平均就是经验W1距离。代码很短import numpy as np def wasserstein_1d(x, y): # 一维经验W1距离排序后配对求平均绝对差 x_sorted np.sort(x) y_sorted np.sort(y) # 若样本数不同需先对齐到相同分位数网格 n min(len(x_sorted), len(y_sorted)) qs np.linspace(0, 1, n) xs np.quantile(x_sorted, qs) ys np.quantile(y_sorted, qs) return np.mean(np.abs(xs - ys))这段代码背后的逻辑就是排序配对。注意一个坑当两组样本数不同的时候不能直接相减必须先量化到相同的分位数网格上再比。我早期偷懒直接算相同长度的样本遇到长度不等就报错或者算错排查了半天才发现是这个原因。一维方法的好处是O(n log n)几百万个点也就毫秒级做实时监控分布漂移完全够用。坏处是它只在一维成立多维不能直接套。4.2 线性规划求解离散分布多维、离散、规模不大的时候老老实实建线性规划。用python-optimaltransport或者直接用scipy的linprog都能做。核心是构造代价矩阵和流守恒约束import numpy as np from scipy.optimize import linprog def wasserstein_lp(xs, ys, px, py): # xs, ys: 源点和目标点坐标 (n,d) 和 (m,d) # px, py: 对应的概率质量各自求和为1 n, m len(xs), len(ys) # 代价矩阵每对点之间的欧氏距离 C np.linalg.norm(xs[:, None, :] - ys[None, :, :], axis2) c C.flatten() # 约束传输矩阵行和px列和py A_eq [] b_eq [] for i in range(n): row np.zeros((n, m)) row[i, :] 1 A_eq.append(row.flatten()) b_eq.append(px[i]) for j in range(m): col np.zeros((n, m)) col[:, j] 1 A_eq.append(col.flatten()) b_eq.append(py[j]) res linprog(c, A_eqnp.array(A_eq), b_eqnp.array(b_eq), bounds(0, None), methodhighs) return res.fun这里把n乘m的传输矩阵拉平成一维决策变量行和列和约束写成等式。数据点几百个的时候这个方法精确又直接一旦上千变量数就是百万级linprog会直接卡死。所以它适合小规模精确验证比如你想校验自己写的近似算法对不对。4.3 Sinkhorn熵正则化加速大规模场景靠的是Sinkhorn迭代。思路是在原始最优传输上加一个熵正则项让问题变得强凸、可以用矩阵缩放交替迭代解决。这引入了一个正则化系数epsilonepsilon越大越快但越偏离精确Wasserstein距离epsilon越小越准但收敛越慢。典型实现用指数核矩阵K然后交替归一化行和列import numpy as np def sinkhorn(a, b, C, eps0.1, n_iter200): # a, b: 边缘分布向量和为1 # C: 代价矩阵 K np.exp(-C / eps) u np.ones_like(a) v np.ones_like(b) for _ in range(n_iter): u a / (K v) v b / (K.T u) transport np.diag(u) K np.diag(v) return np.sum(transport * C)这三个方法我一般这么选一维或近一维直接用排序点数几百以内、要精确值上linprog点数上千、要做优化循环内部调用一定用Sinkhorn。Sinkhorn改成GPU版本后在图像特征对齐里跑几万维的特征完全没问题前提是epsilon调到合适的量级。5. 落地场景从生成模型到分布对齐光会算还不够得知道往哪用。Wasserstein距离的实际应用场景比大多数人想的广我挑几个自己真正做过、有发言权的讲讲。5.1 WGAN的稳定训练与梯度优势最出名的应用就是WGAN。传统GAN用JS散度判别器太强就梯度消失生成器学不动。WGAN把目标换成Wasserstein距离用critic代替判别器因为KR对偶给出的是W1critic不需要输出概率只需要输出一个实数值并且满足Lipschitz约束。我做过一个图像生成的对比实验同样的网络结构JS版本训练几百轮就开始模式崩溃生成样本反复就那几张换成Wasserstein加梯度惩罚之后损失曲线慢而稳地下降生成多样性明显好转。关键参数是梯度惩罚系数通常取10这个数字在多个数据集上被验证比较通用。另外critic的更新次数通常比生成器多常见是每更新一次生成器critic更新五次这样能让W距离估计得更准。5.2 域适应与特征分布对齐域适应是我用Wasserstein距离最多的方向。场景是源域有标签、目标域没标签两个域的特征分布因为采集环境不同而有偏移。做法是在特征提取器后面加一项Wasserstein距离衡量源域特征和目标域特征的分布差异把它作为正则项加到总损失里逼着网络学出跨域一致的特征。这里用Sinkhorn实现最合适因为要在训练循环里反复调用必须快。具体参数上epsilon我一般从0.1开始试太小的epsilon会让Sinkhorn数值不稳定、容易溢出太大的又会让距离估计失真实践中在0.05到0.5之间调。还有个小技巧把特征做一次L2归一化再算距离能显著缓解尺度差异带来的干扰我试过之后对齐效果稳定提升。5.3 形状检索与更宽的应用面跳出深度学习Wasserstein距离在形状分析、点云配准、文档主题分布比较里都有身影。点云配准本质上就是找两个点集之间的最优传输方案把它算出来顺便就得到了最优匹配。文档主题分布之间的比较因为主题是带语义远近的用Wasserstein比用余弦相似度更合理能感知主题A和主题B接近这种结构信息。我还见过有人把它用在推荐系统的分布漂移检测上监控线上用户分布相对训练分布的Wasserstein距离一旦超阈值就触发模型重训。这些场景的共同点都是分布元素之间自带距离结构只要满足这一点Wasserstein距离往往比只看重叠的度量更合适。6. 踩坑记录与参数选择经验最后这部分是我觉得最值钱的内容全是文档里不会写、只有真做过才知道的东西。整理成速查表遇到问题直接对照。6.1 常见问题与排查速查表现象可能原因排查与解决思路Sinkhorn结果全部是NaNepsilon太小导致核矩阵下溢把epsilon调大或对代价矩阵先做归一化训练损失剧烈震荡critic参数更新过猛降低critic学习率增大梯度惩罚系数距离估计值随批量大小变化大经验分布采样噪声增大批量或用移动平均平滑估计多维一维混用报错直接套用了一维排序法多维必须用LP或Sinkhorn计算极慢用了LP求解大规模问题换Sinkhorn并启用GPU距离对尺度异常敏感特征未归一化先做L2或标准化再计算两个几乎相同的分布距离却很大代价矩阵用了错误的度量检查距离定义是否与问题匹配这张表是我一路踩过来的总结。特别是epsilon下溢那条我第一次跑Sinkhorn的时候整个结果全是NaN调了两小时才发现是eps取太小导致exp溢出改成0.1就正常了。提示所有涉及熵正则化参数epsilon的调参都建议先固定其他变量、只扫这一个参数观察距离估计的收敛曲线不要一上来就联合调。6.2 几个反直觉的实操心得第一个心得不是所有分布比较都该用Wasserstein距离。如果你的问题里分布元素之间没有天然的远近关系比如比较两组无序的类别频率那用Wasserstein反而引入虚假的几何结构得不偿失。判断标准很简单问自己元素A和元素B谁离得近这个问题是否有意义有意义才用。第二个心得W1和W2的选择要看下游。W1的对偶形式约束是Lipschitz常数容易用梯度惩罚实现W2和欧氏几何联系紧密有高斯闭式解在统计建模里更方便。我在生成任务里基本用W1在高斯假设成立的特征对齐里用W2分工明确。第三个心得距离数值本身往往不如相对变化有意义。绝对数值受采样、尺度、维度影响很大别去死磕距离应该是多少而要关注这轮比上轮小了多少、线上比训练大多少。我做漂移监控的时候用的是相对基线的比值而不是绝对值鲁棒得多。第四个心得高维下Wasserstein距离会有维度诅咒两个分布的W1估计值随维度增长而膨胀而且从样本估计它的收敛速度会变慢。维度特别高的时候先降维再算或者用切片Wasserstein——把高维分布投影到多个一维方向上各自算一维W1再平均。切片版本计算快、保留了位置敏感性在图像特征这种上千维的场景里我用得很多效果和速度都比直接上高维Sinkhorn更讨喜。我个人在实际操作中的体会是Wasserstein距离最大的价值不在于它是个更高级的度量而在于它把分布差异和几何移动代价这两件事绑在了一起让优化过程有了稳定且符合直觉的方向感。工具没有绝对好坏关键是搞懂它的脾气。希望你把上面的代码和方法拿去跑一遍尤其是自己手算一个一维例子算完再回过头看那个搬土的比喻会有完全不一样的感受。
返回列表