Skip to main content
QUICK REVIEW

[论文解读] Optimal Transport Tools (OTT): A JAX Toolbox for all things Wasserstein

Marco Cuturi, Laetitia Meng-Papaxanthos|arXiv (Cornell University)|Jan 28, 2022
Asphalt Pavement Performance EvaluationEngineering被引用 18
一句话总结

OTT-JAX 是一个基于 JAX 的 Python 工具箱,通过使用熵正则化和低秩近似,实现了高效、可微分的最优传输计算。它支持线性与二次最优传输问题、重心、Gromov-Wasserstein 以及高斯混合匹配,配备可扩展的可微分求解器,适用于机器学习应用。

ABSTRACT

Optimal transport tools (OTT-JAX) is a Python toolbox that can solve optimal transport problems between point clouds and histograms. The toolbox builds on various JAX features, such as automatic and custom reverse mode differentiation, vectorization, just-in-time compilation and accelerators support. The toolbox covers elementary computations, such as the resolution of the regularized OT problem, and more advanced extensions, such as barycenters, Gromov-Wasserstein, low-rank solvers, estimation of convex maps, differentiable generalizations of quantiles and ranks, and approximate OT between Gaussian mixtures. The toolbox code is available at exttt{https://github.com/ott-jax/ott}

研究动机与目标

  • 解决大规模和可微分机器学习应用中最优传输(OT)的计算与可微性挑战。
  • 提供一个统一、高性能的框架,用于求解点云、直方图和测度之间的正则化最优传输问题。
  • 通过 JAX 的自动微分和 JIT 编译实现可微分的最优传输计算,支持深度学习流水线中的端到端训练。
  • 将最优传输能力扩展至标准 Wasserstein 距离之外,包括重心、Gromov-Wasserstein 和软排序操作。
  • 通过低秩近似和几何感知的成本计算实现高效计算,无需显式存储矩阵。

提出的方法

  • 利用 JAX 的自动微分和 JIT 编译,实现高性能的可微分最优传输求解器,适用于 CPU 和 TPU/GPU。
  • 通过 Sinkhorn 算法实现熵正则化,平滑最优传输计划,实现高效、可微分的优化。
  • 引入低秩 Sinkhorn 求解器,通过使用秩-r 因子近似传输矩阵,降低内存和计算成本。
  • 使用几何类隐式计算成本矩阵,避免显式存储——例如,通过核化操作或基于网格的结构计算点云的成本。
  • 通过迭代线性化实现 Gromov-Wasserstein,将二次最优传输问题转化为一系列线性最优传输问题求解。
  • 集成输入凸神经网络(ICNNs),用于学习凸映射,以及通过基于可微分最优传输的重参数化实现软排序操作。

实验结果

研究问题

  • RQ1如何利用现代深度学习框架在大规模场景下高效且可微分地求解最优传输问题?
  • RQ2低秩近似和几何感知计算在多大程度上能够降低最优传输中的内存和时间复杂度?
  • RQ3可微分最优传输能否用于学习结构化表示,如重心或分布之间的映射?
  • RQ4通过最优传输计划的隐式微分如何提升可微分机器学习流水线中的训练稳定性?
  • RQ5Gromov-Wasserstein 和软排序等高级最优传输变体能否实现高效、完全可微分且可扩展的实现?

主要发现

  • OTT-JAX 支持 JAX 的自动微分,实现完全可微分的最优传输,允许通过传输计划进行反向传播。
  • 低秩 Sinkhorn 求解器在不牺牲精度的前提下,显著降低了内存和计算成本,尤其适用于大规模问题。
  • 该工具箱支持重心、Gromov-Wasserstein 距离和软排序数组的端到端可微分计算。
  • 几何类允许隐式计算成本矩阵(例如点云或网格),避免显式存储,实现在网格上 O(dn^{d+1}) 的操作复杂度。
  • 该实现支持通过 Delon 和 Desolneux(2020)提出的可微分近似方法,高效计算高斯混合分布之间的类似 Wasserstein 距离。
  • 该工具箱已具备生产就绪能力,已应用于高级任务,如医学影像数据中的同胚重心计算,以及复杂流形(如螺旋形和瑞士卷)之间的形状匹配。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。