尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

MATLAB实现非标准化Wasserstein距离:从EMD到最优传输实战

MATLAB实现非标准化Wasserstein距离:从EMD到最优传输实战 简介针对一维与二维非标准化Wasserstein-2UW-2距离计算此实现包含基于C与Matlab MEX的完整代码源自Gangbo等人2019年arXiv论文《Unnormalized Optimal Transport》适用于最优输运理论研究、图像配准、分布比较等科研与工程场景适合具备一定最优传输基础和Matlab/C调用能力的研究者。包内共19个文件包括8个C源文件、7个头文件、2个Matlab脚本、1个说明文档和1份许可证压缩包仅31KB结构紧凑按一维与二维求解器分类组织。目前已有373人浏览学习可作为算法复现与二次开发的参考。代码包含完整的一维/二维UW-2求解器类、MEX入口及示例用法C接口与Matlab脚本可对照使用用户可直接在Matlab中调用核心算法也可参考C示例理解求解流程并自行扩展至更高维度或自定义代价函数。1. 项目解读非标准化Wasserstein距离为什么值得单独做一套代码做分布比较、图像检索、点云配准或者生成模型评估的同行应该都绕不开EMDEarth Movers Distance土方移动距离这个名字。它在很多场景下就是Wasserstein距离的离散版本核心思想很朴素把一堆土搬到另一堆土的位置搬运成本最小是多少这个成本就是两个分布之间的距离。但真到自己动手写代码的时候你会发现网上大部分实现都默认处理的是归一化后的概率分布也就是所有样本的权重加起来等于1。可实际工程里我们经常遇到两个分布的总质量本身就不相等比如两张图片的亮度直方图、两组长度不一的传感器读数、两批数量不同的点云数据这时候直接套用标准EMD代码就会出问题算出来的距离既不准确也没有实际物理意义。这个仓库的名字很直白emd代码matlab-unnormalized-optimal-transport它专门解决的就是这个问题——在MATLAB环境下计算非标准化的Wasserstein距离。所谓非标准化就是允许两个输入信号携带的总质量不同代码会把这个质量差也纳入距离度量中而不是先做归一化丢掉这部分信息。这种做法在图像对比、异常检测、分布偏移监控这些场景里非常有用因为质量差异本身往往就是数据变化的重要信号。这个仓库适合谁看如果你在实验室里做信号处理、图像分析或者研究最优传输理论但不想从零推导实现又或者你正在对比两个不完全对齐的数据集合这套MATLAB代码可以直接拿来做基准、跑对比实验省去从论文公式到可运行代码之间的那段痛苦过程。下面我会从算法原理、代码结构、实操步骤和踩坑经验四个维度把整个东西拆开讲清楚。2. 核心算法拆解从Wasserstein距离到非标准化版本2.1 先搞清楚Wasserstein距离到底在算什么要理解非标准化版本得先回到基础定义。Wasserstein距离的数学定义是[ W_p(\mu, u) \left( \inf_{\gamma \in \Gamma(\mu, u)} \int_{X \times X} c(x, y)^p , d\gamma(x, y) \right)^{1/p} ]其中 (\Gamma(\mu, u)) 是所有以 (\mu) 和 ( u) 为边缘分布的联合分布集合。听起来很抽象用人话说就是假设你手里有两堆土(\mu) 表示土堆A的形状( u) 表示土堆B的形状你要规划怎么把A的每一铲土运到B的位置使得总搬运成本最小。(c(x, y)) 就是从位置 (x) 运到位置 (y) 的单份成本通常取欧氏距离。当 (p1) 时这个距离就是EMD不过严格来说经典EMD还要除以总质量做归一化得到的是一个平均搬运成本而不是真正的Wasserstein-1距离。回到MATLAB代码这个话题。标准形式的Wasserstein距离要求 (\mu) 和 ( u) 都是概率测度也就是说两者的总质量都是1。这相当于默认前提是两堆土的总量一样多只是堆放的方式不同。很多实际数据根本满足不了这个假设所以研究者才提出了非标准化的最优传输问题unnormalized optimal transport把质量差异 ( |\mu(X) - u(X)| ) 也放到优化目标里。2.2 非标准化版本做了什么改动这个仓库实现的核心改动体现在目标函数上。标准问题求解的是[ \min_{\gamma} \sum_{i,j} c_{ij} \gamma_{ij} ]约束条件是 (\sum_j \gamma_{ij} \mu_i)、(\sum_i \gamma_{ij} u_j)。而非标准化版本使用的是部分最优传输partial optimal transport的思想或者更常见的做法是引入一个惩罚项处理多余的质量[ \min_{\gamma} \sum_{i,j} c_{ij} \gamma_{ij} \lambda \cdot \text{penalty}(\text{unmatched mass}) ]用人话解释就是允许一部分土可以不搬运但没被搬运的那部分质量要付出代价。这个 (\lambda) 是人为设定的超参数控制着对质量不匹配的容忍程度。(\lambda) 越大算法越倾向于把两边质量强行匹配上(\lambda) 越小算法越容忍质量差的存在。我用一个实际场景来帮助你理解这个设计。假设你在对比两张医学影像的灰度直方图一张是正常组织的扫描结果另一张是病灶区域的扫描结果。两张图像的像素总量可能相同但感兴趣区域的像素占比差异很大。如果用标准Wasserstein距离你必须先把直方图归一化这样就丢失了病灶区域像素更多这个重要信息。用非标准化版本代码会同时计算分布形状的搬运成本和总质量的差异成本给出的距离值更能反映两张图像的真实差异。2.3 MATLAB求解路径的选择逻辑这个仓库选择MATLAB而不是Python有一定历史原因但也不全是。MATLAB自带的优化工具箱非常成熟linprog函数可以高效求解线性规划问题而离散最优传输问题本质上就是一个线性规划。标准形式的Wasserstein距离可以直接转化为[ \min_{x} f^T x \quad \text{s.t.} \quad A_{eq} x b_{eq}, \quad x \geq 0 ]其中 (x) 是传输矩阵展平后的向量(f) 是成本矩阵展平后的向量(A_{eq}) 是边缘约束矩阵(b_{eq}) 是边缘分布值。对于中小规模的数据比如几百个点以内的分布比较linprog的求解速度完全够用而且不需要额外安装第三方库。相比之下Python生态里虽然也有POTPython Optimal Transport这样的成熟库但在实验室环境下如果同事都在用MATLAB做前处理和后分析把EMD计算也放在MATLAB里做能够省掉跨语言数据传递的麻烦。这个仓库的作者显然是沿着这条思路走的——用最顺手的方式解决手头的问题。3. 代码结构梳理与核心函数解读3.1 输入数据的组织方式拿到代码后第一件事是看清楚它期望的输入格式。通常这套代码的核心入口函数需要三个参数两个测度的值向量和一个成本矩阵。值向量就是每个离散位置的权重维度要相同。成本矩阵是 (n \times m) 的矩阵(n) 和 (m) 分别是两个测度支撑点的数量矩阵的每个元素 (C(i,j)) 表示从第一个测度第 (i) 个点运到第二个测度第 (j) 个点的单位成本。我在实际操作中遇到的最常见的错误是维度没对齐。比如说你的第一个分布是在 0 到 100 的区间上均匀采样了 100 个点第二个分布是在 0 到 80 的区间上采样了 80 个点那么成本矩阵必须是 (100 \times 80)不能是方阵。如果你用的是二维点云成本矩阵就是所有点对之间的欧氏距离矩阵可以用pdist2快速计算% 假设 pointSet1 是 n×d 矩阵pointSet2 是 m×d 矩阵 C pdist2(pointSet1, pointSet2);这段代码非常简单但它是整个流程的地基成本矩阵如果算错了后面所有距离值都是空中楼阁。3.2 线性规划模型的构造细节接下来是核心的建模步骤。假设测度1有 (n) 个支撑点测度2有 (m) 个支撑点传输变量 (x) 是一个长度为 (n \times m) 的向量表示从每个测度1的点运到每个测度2的点的质量。目标函数的系数向量 (f) 就是把成本矩阵按列或按行展平。两种展平方式都可以但要保证和约束矩阵的构造方式一致这是一个最容易写错的地方。约束矩阵 (A_{eq}) 的结构是 ((nm) \times (n \times m))。前 (n) 行对应测度1的每个点这一行只在和该点相关的 (m) 个变量位置为1约束值是该点携带的质量 (\mu_i)。后 (m) 行对应测度2的每个点同样在该点相关的 (n) 个变量位置为1约束值是 (u_j)。这样就能保证求出来的传输矩阵它的行和等于 (\mu)、列和等于 ( u)。非标准化版本的改动就在于最后一步如果总质量不相等直接构造严格等式约束会无解因为 (\sum \mu_i eq \sum u_j)。所以代码里会额外引入松弛变量或者把约束从等于改成小于等于再在目标函数里对松弛部分加惩罚。这个仓库具体选择了哪一种做法你需要看一下主函数的实现但我更推荐用的是部分最优传输的松弛形式因为它有明确的理论性质不会因为惩罚系数设置不当导致结果变得很怪。3.3 调用求解器与结果解析MATLAB的linprog函数在这套代码里承担了核心求解任务。典型调用方式是options optimoptions(linprog, Algorithm, dual-simplex, Display, off); [x, distance, exitflag] linprog(f, [], [], Aeq, beq, lb, [], [], options);其中lb是长度为 (n \times m) 的全零向量对应传输质量的非负约束。exitflag非常关键它告诉你求解是否成功。如果exitflag等于 1说明找到了最优解如果是负数说明遇到了数值问题或者约束冲突。得到最优解 (x) 之后把它重新变形成 (n \times m) 的矩阵Gamma reshape(x, n, m);Gamma(i,j)就是从测度1的第 (i) 个点运到测度2的第 (j) 个点的质量。而最终的Wasserstein距离值就是所有 (C(i,j) \times \Gamma(i,j)) 的和emd_value sum(sum(C .* Gamma));注意如果是非标准化版本最终的距离值还应该加上质量差惩罚项你需要看代码里是否已经把这一部分计入返回值还是需要你额外处理。4. 实操过程记录从零跑通一次距离计算4.1 第一步准备模拟数据我建议你第一次跑这套代码时先用一组简单的模拟数据验证流程而不是直接上真实数据。比如在MATLAB命令行里构造两个高斯分布% 测度1均值0方差1在-5到5之间采样50个点 x1 linspace(-5, 5, 50); mu1 exp(-0.5 * x1.^2) / sqrt(2*pi); % 测度2均值1方差1同样在-5到5之间采样50个点 x2 linspace(-5, 5, 50); mu2 exp(-0.5 * (x2-1).^2) / sqrt(2*pi); % 成本矩阵取两两欧氏距离的平方或直接取距离取决于你的应用 C abs(x1 - x2); % 一维情况下的距离矩阵这段代码产生的两个分布形状相同、位置不同理论上它们的Wasserstein距离应该是1左右均值平移量。用这套代码算一下如果结果接近1说明流程是通的。如果你把mu2整体乘以2得到的总质量就是测度1的两倍这时候标准EMD和非标准化EMD的结果会有明显差异你可以直观感受到这套代码的特殊价值。4.2 第二步运行核心函数并检查结果把数据喂给核心函数后建议做三件事。第一件事是检查exitflag确认线性规划求解成功。第二件事是检查传输矩阵Gamma的结构合理性——大部分质量应该集中在靠近对角线附近的位置这符合就近搬运的直觉。如果发现质量分散在距离很远的位置说明成本矩阵构造有误或者数据维度出了问题。第三件事是验证约束是否满足分别计算sum(Gamma, 2)和sum(Gamma, 1)对比原始的mu1和mu2误差应该在 (10^{-6}) 量级以内。实测下来这套MATLAB代码在数据量达到几百个支撑点时运行速度还是可以接受的。比如有200个点和200个点的比较线性规划的变量数是40000dual-simplex算法通常能在几秒内求解完毕。但如果你的数据点数量上升到几千求解时间会急剧膨胀这时候你应该考虑换用熵正则化的Sinkhorn算法。有些仓库的代码里也附带了Sinkhorn的加速实现专门应对大规模数据这是很实用的设计值得你仔细看看。4.3 第三步批量计算的效率优化技巧实际项目中很少只算一对分布的距离更多时候是拿一个查询分布和数据库里成千上万个分布做对比。这时候逐个调用linprog会非常慢。我的建议是先在核心函数外面套一层缓存机制如果成本矩阵固定不变就提前把它算好存下来不要每次调用都重新计算。另外linprog每次调用都要重新构建稀疏矩阵这个开销在循环里会被放大。更好做法是提前把约束矩阵的稀疏结构建好循环里只更新beq向量% 预先分配稀疏矩阵 Aeq注意这里的索引要提前计算好 % 每次循环只需更新 beq [mu; nu];这样整体速度能提升好几倍尤其是当你有几百个查询样本需要批量处理时这个优化不是小打小闹而是能决定你今晚能不能按时收工的关键。5. 常见问题与调试经验速查5.1 维度不匹配与内存爆炸这套代码最常见的报错是约束矩阵的维度对不上。比如测度1有100个点测度2有80个点那么约束矩阵应该是 (180 \times 8000)。如果你在改代码时不小心把某个维度写成了方阵linprog会直接报错或者算出一个完全错误的结果。这类问题在所有MATLAB最优传输实现里都很常见因为维度关系不够直观。另一个是内存问题。当两个测度各有1万个点时(A_{eq}) 矩阵的大小是 (20000 \times 10^8)即使用稀疏矩阵存储也相当吃内存。我的建议是遇到这种情况先对数据做聚类或抽稀把支撑点数量压到2000以内再计算距离。因为Wasserstein距离对支撑点的数量有一定鲁棒性只要采样方式一致距离值的相对关系基本保持不变批量筛选场景下已经足够使用。5.2 数值稳定性与惩罚系数的调节在非标准化版本中惩罚系数 (\lambda) 的取值对结果影响很大。我在调试时发现如果 (\lambda) 设得太小算法倾向于完全忽略质量差回退成形状匹配非标准化就失去意义了如果设得太大数值上容易导致线性规划的条件数变差求解器可能报numerical issues的警告。一个可行的经验值是先算一下成本矩阵的平均值然后把 (\lambda) 设在这个平均成本的量级比如 (\lambda \text{mean}(C(:)))然后上下调整几个数量级做敏感性分析。具体到代码调试时我习惯写一个简单的循环把 (\lambda) 从 (0.01 \times \text{mean}(C)) 扫到 (100 \times \text{mean}(C))观察距离值的变化曲线。如果曲线在某个区间内基本稳定说明这个区间是合理的工作区间如果在某个点突然跳变那大概率是数值不稳定需要重新审视成本矩阵的构造或者数据预处理方式。5.3 加权测度的空值处理真实数据里经常出现某些位置的权重为0的情况也就是测度在某些支撑点上没有质量。这本身没问题但要注意不能让所有质量都为0否则约束里会出现全零行导致问题退化。我在处理传感器数据时遇到过这种情况某个时间段传感器没有采集到信号对应位置的权重就是0。解决办法很简单在调用函数之前把权重全为0的支撑点剔除掉同时同步处理成本矩阵的行列这样既减小了问题规模也避免了退化问题。5.4 MATLAB工具箱依赖判断这套代码依赖MATLAB的Optimization Toolbox因为核心求解器是linprog。我见过有些同事的电脑上只有基础MATLAB环境没有安装工具箱一跑就报Undefined function linprog。遇到这种情况有两个出路第一个是安装Optimization Toolbox这是最省事的方案第二个是改用MATLAB自带的fmincon或者其他替代函数但要自己做线性规划建模会比较绕。如果你在实验室部署这套代码给其他人用建议先检查目标机器的工具箱情况省得到时候手忙脚乱。6. 扩展思考这套代码还能往哪些方向延伸这套代码给我最大的启发是非标准化这个设计思路的应用范围远超出它本身的实现。举个例子在图像检索任务中传统方法提取颜色直方图后直接计算直方图的欧氏距离或者卡方距离对光照变化非常敏感。改用非标准化Wasserstein距离后因为引入了质量差异项算法能同时感知到颜色分布变化和总体亮度变化两个维度在光照条件多变的数据集上效果通常更好。另外一个值得尝试的方向是把这套MATLAB代码和深度学习框架结合。虽然现在大多数研究工作都在Python生态里做但工程落地阶段用MATLAB做数据处理、用深度学习框架做特征提取、再用这套代码做最终评估是一条很实际的组合路径。尤其是在需要编写论文实验部分的场景下MATLAB绘图和数据分析的便利性依然无人能替代。如果你对非标准化最优传输的数学理论感兴趣建议从三个方向深入部分最优传输理论、不平衡最优传输unbalanced optimal transport、以及熵正则化的非标准化版本。理论吃透了再看这套代码你会发现每个函数的取舍背后都有明确的数学动机而不仅仅是工程上的妥协。我自己在理论和代码之间来回印证了好几遍才真正理解为什么松弛变量要那样加、惩罚项为什么要那样设置。这种理解程度对于你后续修改代码适配自己的数据是不可或缺的。本文还有配套的精品资源点击获取
返回列表