
1. 项目概述从“搬箱子”到“量宇宙”——最优传输理论的工程化解读最近在整理一些关于度量学习和分布匹配的笔记发现“最优传输”这个概念出现的频率越来越高。从最初在计算机图形学里搞纹理合成到后来在机器学习里做领域自适应、生成模型甚至是在单细胞测序数据分析里比对不同样本的细胞分布Optimal Transport都像个“万能胶”一样把看似不相关的问题给统一到了一个优雅的数学框架下。但说实话很多材料要么一上来就是一堆测度论和凸分析的公式让人望而却步要么就只给个调用Python库POT的代码片段知其然不知其所以然。我自己也是从“这玩意儿到底在算什么”的困惑中走过来的。今天我就想抛开那些复杂的数学外衣用最“工程化”的视角把最优传输的核心——Wasserstein距离、Gromov-Wasserstein距离以及它们的融合变体Fused Gromov-Wasserstein距离——给拆解明白。你可以把这篇文章看作是一个“用户手册”核心目标就一个当你手里有两堆数据点比如两张图片的像素分布、两个数据集的特征、两个生物样本的细胞群想知道它们到底有多“像”或者想找到它们之间点对点的“最佳配对”时你知道该抄起哪把“尺子”来量以及这把尺子背后的“拧螺丝”的细节是什么。2. 核心思想拆解最优传输到底在解决什么问题2.1 从“搬沙子”到“度量距离”最优传输最经典的例子是“搬沙子”。想象你有两堆沙子分布在不同位置你的任务是用最小的总力气把第一堆沙子的形状搬成第二堆的样子。这里的“力气”通常定义为移动距离的某种函数比如平方。最优传输就是要找到那个“搬法”使得总成本最低。这个最低总成本就是一种距离它衡量了把一堆沙子变成另一堆的“难易程度”。在数学和工程上这两堆“沙子”就是两个概率分布。我们有很多方法来衡量两个分布的距离比如大家熟悉的KL散度、JS散度。但这些散度有个问题它们只关心分布函数本身在每个点的值不关心点的“位置”关系。举个例子两个正态分布一个均值是0一个是10它们的概率密度函数在图像上几乎没有重叠区域KL散度会趋于无穷大告诉你它们“完全不同”。但这显然忽略了它们只是整体偏移了一下的事实。而Wasserstein距离则不同它通过计算“搬动”质量所需的成本自然地考虑了分布支撑集数据点所在空间的几何信息。均值0和均值10的两个分布Wasserstein距离会给出一个大约为10的有限值直观地反映了“需要移动10个单位”这个事实。2.2 为什么是Wasserstein关键优势剖析Wasserstein距离特别是p-Wasserstein距离之所以在机器学习社区大火是因为它解决了传统方法的几个痛点对低概率区域更鲁棒KL散度在分母概率为0时会爆炸对估计误差极其敏感。Wasserstein距离基于累积分布更加平滑稳定。能反映几何结构如前所述它考虑了数据点所在空间的距离度量。这使得它特别适合衡量生成模型如GAN生成图像的质量。一个模糊但位置正确的小狗图片和一个清晰但位置错乱的小狗图片在像素级的L2距离上可能后者更差但从分布感知的Wasserstein距离看前者可能更“像”真实的小狗分布。提供自然的插值最优传输计划本身就给出了一个点对点的映射。这意味着你不仅知道两个分布有多远还能知道如何将一个分布“渐变”到另一个分布这被称为“Wasserstein重心”或“传输插值”在图像变形、风格迁移中非常有用。然而标准的Wasserstein距离要求两个分布定义在同一个度量空间里。也就是说你的两堆“沙子”必须撒在同一个“操场”上我们才能计算从A点搬到B点的距离。但现实问题往往更复杂。2.3 当“操场”不同时Gromov-Wasserstein的登场设想两个场景一是比较两个不同社交网络的结构二是对齐两个不同语言语义空间中的词向量。在这两个例子里我们想比较的是两个分布但它们各自内部的点网络中的用户、语义空间中的词所处的“空间”根本不同。我们无法直接计算用户A到词“苹果”的欧氏距离。这时就需要Gromov-Wasserstein距离。它的核心思想非常巧妙我们不比较点与点之间的绝对位置而是比较点与点之间的相对关系。具体来说它比较的是两个度量空间各自内部的距离或相似性结构。对于第一个空间中的任意两点计算它们的距离在第二个空间中找到对应的两点也计算它们的距离然后看这两对距离的差异。GW距离寻找的是一种“对应关系”耦合使得从一个空间内部距离矩阵到另一个空间内部距离矩阵的“失真”最小。这就好比比较两幅拼图。Wasserstein距离要求两幅拼图碎片形状完全一样同一空间只是位置不同。而Gromov-Wasserstein距离则允许碎片形状不同不同空间但它关心的是碎片之间的“连接关系”是否相似——比如某块碎片在所有碎片中是更靠近中心还是更靠近边缘它和周围几块碎片的相对位置关系是怎样的。通过匹配这种相对结构我们就能判断两幅拼图描绘的是否是同一个场景。2.4 融合的力量Fused Gromov-Wasserstein距离那么有没有一种方法能同时利用绝对位置信息和相对结构信息呢这就是Fused Gromov-Wasserstein距离要解决的问题。FGW距离可以看作是在Wasserstein和Gromov-Wasserstein之间取了一个权衡。它定义了一个融合的成本函数将点对之间的特征距离Wasserstein部分和关系结构距离GW部分线性结合起来。公式上它寻找一个传输计划最小化以下两项的加权和将第一个空间点i的特征搬到第二个空间点j的特征所需的成本如欧氏距离。第一个空间中点i与其他点k的关系与第二个空间中点j与其他点l的关系之间的差异成本。通过一个参数α在0到1之间来控制两者的权重。当α1时FGW退化为只关心特征的Wasserstein距离当α0时退化为只关心结构的Gromov-Wasserstein距离当α取中间值时则同时考虑两者。这非常实用。例如在图形匹配中每个节点既有自身的特征如用户的年龄、兴趣又有与其他节点的关系如好友链接。FGW距离允许我们同时匹配节点的属性和图的结构。在单细胞数据分析中细胞既有基因表达特征又有在发育轨迹或空间转录组中的位置关系FGW可以同时对齐这两个层面的信息。3. 核心算法与实操要点3.1 Wasserstein距离的计算从线性规划到Sinkhorn迭代理论上计算两个离散分布比如两个直方图之间的Wasserstein距离是一个线性规划问题。假设第一个分布有n个点权重为p第二个有m个点权重为q。我们需要找到一个n×m的非负传输矩阵T其行和等于p列和等于q。目标是最小化总成本 ∑ᵢ∑ⱼ Tᵢⱼ * Cᵢⱼ其中Cᵢⱼ是点i到点j的成本如距离的p次方。直接解这个线性规划复杂度是O(n³ log n)级别的对于大规模数据如图像n为像素数完全不可行。实操中的救星熵正则化与Sinkhorn算法为了加速计算Cuturi在2013年提出了一个里程碑式的想法在目标函数中加入传输矩阵T的熵正则项 -ε H(T)。修改后的问题变成了一个严格凸的问题并且其解具有特定的形式T diag(u) * K * diag(v)其中K exp(-C/ε)而u和v是两个正的对角缩放向量。求解u和v可以通过一种极其简单的迭代算法——Sinkhorn迭代也称为矩阵缩放算法来完成初始化 v 1全1向量。交替更新u p ./ (K v) v q ./ (Kᵀ u)。./ 表示逐元素除法重复直到收敛。这个算法的复杂度每次迭代是O(nm)如果利用距离矩阵C的结构如网格距离甚至可以降到O(n log n)。参数ε控制正则化的强度ε越大解越平滑熵越大计算越快但距离估计有偏ε越小越接近原始Wasserstein距离但计算越不稳定。实操心得通常从ε0.1左右开始尝试观察结果对ε的敏感性。对于只是用来做损失函数训练神经网络的情况较大的ε如1.0往往能提供更稳定的梯度。注意Sinkhorn迭代涉及指数运算exp(-C/ε)当C/ε很大时可能导致数值下溢。稳定的实现通常会先减去C矩阵的最大值即计算K np.exp((-C np.max(C)) / epsilon)并在迭代中处理可能的数值问题。3.2 Gromov-Wasserstein距离的计算四维张量与交替优化GW距离的计算比W更复杂。它的成本函数涉及四维张量L(C_s, C_t, i, j, k, l) |C_s[i, k] - C_t[j, l]|²其中C_s和C_t分别是源空间和目标空间内部的距离平方矩阵。离散GW的优化问题同样是寻找耦合矩阵T最小化 ∑ᵢⱼ∑ₖₗ Tᵢⱼ Tₖₗ L(C_s, C_t, i, j, k, l)。这是一个非凸的二次规划问题。主流求解方法条件梯度法Frank-Wolfe与熵正则化版本条件梯度法在每一步计算当前耦合T下的梯度一个n×m矩阵然后求解一个线性规划问题找到下降方向通常是一个排列矩阵即一个点只对应另一个点再沿着这个方向进行线搜索更新T。这个方法能保证收敛到局部极小并且迭代中产生的解是稀疏的很多0元素。熵正则化GWEGW同样可以引入熵正则项将问题转化为一个可以通过类似Sinkhorn的迭代来求解的问题但此时每次迭代需要计算一个与梯度相关的核矩阵计算复杂度为O(n²m nm²)。对于大规模问题需要采用采样或低秩近似等方法。实操要点距离矩阵预处理计算GW前务必对输入的距离矩阵C_s和C_t进行标准化例如除以矩阵的中位数或最大值。因为GW损失函数比较的是距离的差异如果两个空间的距离尺度差异巨大一个空间距离在0-1另一个在0-1000优化会不稳定结果会偏向于匹配距离尺度而非结构。初始化的重要性GW问题非凸初始化很重要。常见的策略包括使用均匀分布初始化T或者先用一些快速算法如基于特征向量的谱方法得到一个粗糙的对应作为“热启动”。对称性问题如果要计算两个分布之间的GW距离而不是做匹配结果应该是对称的。但迭代算法可能因为初始化不同而收敛到不同的局部极小导致不对称。一个实践技巧是计算两次A到BB到A然后取平均值或最小值。3.3 Fused Gromov-Wasserstein距离的计算融合成本的优化FGW的优化问题结合了W和GW的成本其目标函数为 (1-α) * ∑ᵢ∑ⱼ Tᵢⱼ * C_{feature}(i, j) α * ∑ᵢⱼ∑ₖₗ Tᵢⱼ Tₖₗ L(C_s, C_t, i, j, k, l) 其中C_{feature}是特征空间的距离矩阵。求解算法目前最常用的也是基于条件梯度法FGW-FW或熵正则化FGW-EGW。其步骤与GW类似但在计算梯度时需要同时考虑特征成本项和结构成本项。对于特征成本部分梯度就是C_{feature}。对于结构成本部分梯度是一个矩阵G其中G[i,j] α * ∑ₖₗ L(C_s, C_t, i, j, k, l) * T[k, l]。总的梯度是 (1-α)*C_{feature} G。 然后同样通过解线性规划或熵正则化迭代来更新耦合矩阵T。参数α的选择这是FGW应用中的关键超参数。无监督场景如果没有标签信息指导通常采用交叉验证或基于任务性能的网格搜索。一个经验法则是如果特征信息非常可靠且判别性强α设小一些如0.1-0.3如果结构信息更可靠如图形拓扑α设大一些如0.7-0.9。半监督/有监督场景如果有一小部分已知的对应点对可以将α作为一个可学习参数与耦合矩阵T一起优化使得已知对应点对的匹配概率最大化。多尺度尝试在实际项目中我通常会绘制一个“α-任务性能”曲线例如α从0到1步长0.1观察下游分类或聚类精度的变化选择性能平台区的中点作为最终参数这样结果对参数微调不敏感更鲁棒。4. 工程应用场景与代码实践4.1 场景一领域自适应与样本对齐假设我们有一个带标签的源域数据集如清晰照片和一个无标签的目标域数据集如雾天照片分布不同但任务相同如分类。我们希望利用源域模型帮助目标域。传统方法最小化源域和目标域特征间的MMD或对抗损失。OT方法计算源域和目标域特征分布之间的Wasserstein或FGW距离并将其作为正则项加入损失函数。最优传输矩阵T提供了源域样本和目标域样本之间的软对应关系。我们可以利用T来特征对齐将源域特征“搬运”到目标域分布附近。标签传播通过T将源域样本的标签以加权方式传递给目标域样本为目标域生成伪标签。import numpy as np import ot # Python Optimal Transport library # 假设 source_feats, target_feats 分别是源域和目标域的特征矩阵 # source_labels 是源域标签 n_source source_feats.shape[0] n_target target_feats.shape[0] # 1. 计算特征成本矩阵 C_feat ot.dist(source_feats, target_feats) # 默认欧氏距离平方 # 2. 计算或定义结构成本矩阵例如基于特征空间k近邻图 # 这里以简单的特征距离作为结构信息的替代仅作示例 C_source_struct ot.dist(source_feats, source_feats) C_target_struct ot.dist(target_feats, target_feats) # 3. 定义均匀分布假设样本权重相同 p np.ones(n_source) / n_source q np.ones(n_target) / n_target # 4. 计算FGW耦合矩阵 alpha 0.5 # 平衡参数 T_fgw, log ot.gromov.fused_gromov_wasserstein( MC_feat, C1C_source_struct, C2C_target_struct, pp, qq, loss_funsquare_loss, alphaalpha, verboseTrue, logTrue ) # 5. 标签传播为目标域生成软标签 # T_fgw[i, j] 表示源域样本i与目标域样本j的对应强度 target_soft_labels T_fgw.T source_labels_one_hot # 假设source_labels_one_hot是one-hot编码 target_pseudo_labels np.argmax(target_soft_labels, axis1)注意事项直接使用所有样本计算OT复杂度很高。对于大规模数据常用做法是使用小批量mini-batchOT。先对特征进行降维如PCA。使用近似算法如基于子采样的SW切片Wasserstein距离。4.2 场景二图形与点云匹配给定两个图G1和G2节点带有特征边带有权重或距离。目标是找到节点之间的对应关系。传统方法谱匹配、图神经网络。OT方法将每个图表示为一个分布每个节点是一个“质点”其权重可以相同也可以由节点度中心性等决定。特征成本矩阵使用节点特征的距离如属性向量欧氏距离。结构成本矩阵使用图内部的节点距离矩阵如最短路径长度、扩散距离。然后计算FGW距离及其耦合矩阵TT[i,j]的值就给出了节点i来自G1与节点j来自G2的匹配概率。import networkx as nx import ot # 构建两个图 G1, G2 G1 nx.erdos_renyi_graph(20, 0.3) G2 nx.erdos_renyi_graph(25, 0.3) # 为节点添加随机特征 for i in G1.nodes(): G1.nodes[i][feat] np.random.randn(5) for i in G2.nodes(): G2.nodes[i][feat] np.random.randn(5) # 提取节点特征矩阵和距离矩阵 feats1 np.array([G1.nodes[i][feat] for i in G1.nodes()]) feats2 np.array([G2.nodes[i][feat] for i in G2.nodes()]) C_feat ot.dist(feats1, feats2) # 计算图内距离矩阵这里用最短路径长度对于不连通图需处理无穷大 def graph_distance_matrix(G): n G.number_of_nodes() D np.zeros((n, n)) lengths dict(nx.all_pairs_shortest_path_length(G)) for i in range(n): for j in range(n): D[i, j] lengths[i].get(j, float(inf)) # 不连通则距离为无穷大 # 处理无穷大用一个大的数值代替例如最大有限距离的倍数 max_finite np.max(D[D ! float(inf)]) D[D float(inf)] 2 * max_finite if max_finite 0 else 10 return D C_struct1 graph_distance_matrix(G1) C_struct2 graph_distance_matrix(G2) # 计算FGW p ot.unif(feats1.shape[0]) q ot.unif(feats2.shape[0]) alpha 0.7 # 更侧重结构 T_fgw ot.gromov.fused_gromov_wasserstein( MC_feat, C1C_struct1, C2C_struct2, pp, qq, loss_funsquare_loss, alphaalpha ) # 找到最可能的匹配 matching np.argmax(T_fgw, axis1) # 对于G1中的每个节点找到G2中对应概率最高的节点实操心得图距离矩阵的计算可能很耗时尤其对于大图。在实际中常常使用更高效的结构描述符如节点的谱特征拉普拉斯矩阵的特征向量。节点的子图统计量如不同半径的邻居节点度分布。使用图神经网络学习到的节点嵌入作为特征此时结构信息已隐含在嵌入中可能只需计算Wasserstein距离即可。4.3 场景三单细胞多组学数据整合这是计算生物学中的一个热门应用。我们有两个单细胞数据集分别测量了同一批细胞的不同模态信息如转录组和染色质可及性或者同一模态但来自不同批次、不同个体的数据。目标是整合这些数据找到不同数据集之间细胞的对应关系。挑战不同模态的数据特征空间完全不同基因表达值 vs. 染色质开放区域信号直接计算特征距离无意义。FGW解决方案将每个数据集表示为一个分布每个细胞是一个点。特征成本可以设置为0矩阵如果特征不可比或者使用一些先验知识构建例如基于基因-调控区域的关联信息。更常见的做法是先各自降维如PCAUMAP在低维嵌入空间计算特征距离这个低维空间被认为捕捉了生物状态的连续变化。结构成本使用细胞在各自数据空间中的相似性矩阵。这可以通过构建细胞-细胞k近邻图然后计算图扩散距离或最短路径距离来获得。这个结构反映了细胞之间的发育轨迹或状态连续性。计算FGW距离和耦合矩阵。耦合矩阵T给出了跨数据集细胞的对应关系可用于批次校正、多模态数据对齐和标签转移。注意事项单细胞数据通常维度极高数万个基因必须进行特征选择和降维。结构矩阵细胞相似性的质量至关重要。噪声大的数据会导致相似性计算不准严重影响匹配结果。预处理步骤去噪、归一化、高变基因选择需要仔细进行。由于细胞数量可能很大数万至数十万直接计算OT不可行。需要采用近似方法如基于子采样的重心估计、小批量OT或者使用可微分OT层嵌入到神经网络中通过梯度下降间接学习对齐。5. 常见陷阱、调试技巧与性能优化5.1 数值稳定性问题Sinkhorn迭代溢出计算K exp(-C/ε)时如果C/ε值过大exp结果可能下溢为0导致后续除法出现NaN。解决方案使用ot.bregman.sinkhorn或ot.sinkhorn等库函数它们内部实现了数值稳定的log-domain Sinkhorn算法。如果自己实现可以采用log_u和log_v的迭代形式在log空间进行计算。GW/FGW中的距离矩阵距离矩阵中的值如果过大或过小会导致优化问题尺度不佳。务必进行标准化。我通常采用C C / np.median(C)或C C / np.max(C)。确保两个空间的距离矩阵尺度大致相当。非对称距离矩阵如果输入的距离矩阵不是对称的由于数值误差或计算方法导致GW/FGW计算可能会出错。使用C (C C.T) / 2确保对称性。5.2 计算复杂度与加速策略OT的计算复杂度是主要的应用瓶颈。对于大规模Wasserstein距离切片Wasserstein距离通过随机投影将高维分布投影到一维计算一维Wasserstein距离有解析解然后对多个随机方向取平均。这是最常用的近似方法之一复杂度降至O(n log n)。小批量Sinkhorn在深度学习训练中每次只从数据集中采样一个小批量进行计算。需要仔细设计采样策略以保证mini-batch能代表整体分布。分层/多尺度OT先将数据聚类或粗化在粗粒度上计算OT再将结果细化和传播到细粒度。对于大规模GW/FGW距离低秩近似如果距离矩阵C是低秩的或可以近似为低秩可以大幅降低GW中四维张量计算的开销。采样方法只随机采样一部分点对来计算GW损失是一种随机优化策略。使用GPU加速Sinkhorn迭代和矩阵运算非常适合GPU并行。POT库部分支持GPUGeomLoss库则提供了PyTorch/CUDA后端的高效实现。提前终止在迭代优化中如果损失函数下降已非常缓慢可以提前终止不必追求绝对收敛。5.3 超参数选择经验熵正则化系数 ε在Sinkhorn中ε越大计算越快、越稳定但偏差越大。一个实用的策略是使用一个“退火”计划开始时用较大的ε快速得到一个粗略的对齐然后逐渐减小ε进行微调。FGW平衡参数 α如前所述绘制性能曲线是关键。另一个技巧是使用“自适应的α”先分别用纯Wα1和纯GWα0计算耦合矩阵观察匹配质量。如果其中一个明显不合理则可以确定α应该更偏向另一方。损失函数选择GW/FGW中的loss_fun参数常用square_lossL2损失或kl_loss基于KL散度的损失。square_loss更常用但对异常值更敏感kl_loss更鲁棒但计算可能稍慢。根据数据噪声情况选择。5.4 结果评估与可视化OT的输出是一个耦合矩阵T如何评估其好坏对于有真实匹配的任务如图形匹配、单词翻译计算匹配精度将T每行取argmax得到硬匹配与真实匹配对比。计算最近邻召回率对于源域每个点其真实匹配点在目标域中是否出现在高概率对应的前k个点中。对于无监督任务如分布对齐、生成模型下游任务性能这是黄金标准。将对齐后的数据用于分类、聚类等任务看性能是否提升。可视化这是非常有效的调试手段。绘制耦合矩阵T的热图。一个好的对齐热图应该近似一个块对角矩阵如果数据有类别结构或一个相对集中、清晰的模式。使用t-SNE或UMAP将两个域的数据投影到同一二维空间用颜色区分域。在对齐前两个域的点应该分开在对齐后相同类别的点应该混合在一起。对于FGW可以分别绘制特征成本矩阵和结构成本矩阵的热图并与最终的耦合矩阵对比理解算法是如何权衡两者的。6. 高级话题与未来展望6.1 可微分最优传输与深度学习将OT层嵌入神经网络使其能够端到端训练是当前研究的热点。核心思想是让Sinkhorn迭代或GW优化的过程可微分允许梯度反向传播。这带来了两大好处学习特征表示OT距离作为损失函数可以指导神经网络学习使得两个分布更接近的特征表示。例如在无监督领域自适应中网络的特征提取器会被训练以最小化源域和目标域特征之间的Wasserstein距离。学习成本矩阵在FGW中特征成本矩阵C_feat通常由特征向量的距离决定。在可微分框架下我们可以进一步引入一个神经网络来学习这个成本矩阵或者学习特征本身使得最终的匹配更符合任务需求。这相当于让模型自己去发现“什么特征和什么结构对于匹配是重要的”。实现上可以使用ott、GeomLoss或PyTorch自定义函数利用隐函数求导来实现。需要注意的是OT迭代的数值稳定性在反向传播中要求更高。6.2 不平衡最优传输经典OT要求两个分布的总质量必须相等行和列和为1。但在许多实际应用中比如目标检测一张图片中目标的数量可变或部分匹配一个图形是另一个图形的子图质量可能不相等。不平衡最优传输通过放松边际约束允许在支付一定惩罚成本的前提下创造或销毁质量。这通过在原目标函数中加入对行和、列和偏离的惩罚项来实现。库函数如ot.sinkhorn_unbalanced提供了相关实现。6.3 计算拓扑与OT的结合这是一个新兴的前沿方向。持久同调等计算拓扑工具可以描述数据的拓扑特征如连通分支、空洞。将数据的拓扑特征持久图本身视为一个分布然后计算它们之间的Wasserstein距离即所谓的“持久图Wasserstein距离”为比较数据的拓扑结构提供了强大的工具。而Gromov-Wasserstein距离则可以用来比较不同持久图之间的结构关系或者将拓扑特征与几何特征进行融合分析。从我个人的实践经验来看最优传输从一个精妙的数学理论发展到今天在各个工程领域落地的实用工具其核心魅力在于它提供了一种基于几何和概率的、自然的、可解释的数据比较和转换视角。它不像一些黑箱模型OT的耦合矩阵是白箱的你可以清晰地看到每一个数据点是如何被“搬运”和匹配的这为结果分析提供了极大的便利。当然它的计算成本依然是阻碍其大规模应用的拦路虎但随着近似算法、硬件加速和可微分框架的发展这个障碍正在被迅速打破。对于任何需要处理分布比较、数据对齐或结构化预测问题的工程师来说最优传输工具箱里Wasserstein、Gromov-Wasserstein和Fused Gromov-Wasserstein这几把“尺子”都值得花时间深入理解并熟练使用。