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 ツールボックスです。線形および二次的 OT 問題、バーチャル・バーチャル、グロモフ・ワッサーシュタイン、およびガウス・ミックスチャージ・マッチングをサポートし、スケーラブルなソルバーと機械学習応用のための自動微分を備えています。

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)の計算的および微分可能性の課題を解決すること。
  • ポイントクラウド、ヒストограм、測度間の正則化された OT 問題を解くための統合的かつ高性能なフレームワークを提供すること。
  • JAX の自動微分と JIT コンパイルを活用した微分可能な OT 計算を可能にし、ディープラーニングパイプラインにおけるエンド・ツー・エンドの学習を支援すること。
  • 標準的なワッサーシュタイン距離を超えて、バーチャル・バーチャル、グロモフ・ワッサーシュタイン、およびソフト・ソーティング操作を含む OT 機能を拡張すること。
  • 低ランク近似と幾何学に配慮したコスト計算により、明示的な行列格納を回避し、効率的な計算を実現すること。

提案手法

  • JAX の自動微分と JIT コンパイルを活用し、CPU や TPU/GPU での高性能な微分可能な OT ソルバーを実現すること。
  • シンコールン法によるエントロピー正則化を実装し、最適輸送計画を滑らかにし、効率的で微分可能な最適化を可能にすること。
  • 低ランクシンコールン・ソルバーを導入し、輸送行列をランク-r 要素で近似することで、メモリと計算コストを削減すること。
  • 幾何クラスを用いてコスト行列を暗黙的に計算し、明示的な格納を回避する(例:カーネル化された操作によるポイントクラウドやグリッド構造)。
  • 反復的線形化を用いてグロモフ・ワッサーシュタインをサポートし、二次的 OT 問題を一連の線形 OT 問題に還元すること。
  • 入力凸ニューラルネットワーク(ICNN)を統合し、凸写像と微分可能な OT ベースの再パラメータ化によるソフト・ソーティング操作を学習可能にすること。

実験結果

リサーチクエスチョン

  • RQ1現代のディープラーニングフレームワークを用いて、大規模な状況下で効率的かつ微分可能な最適輸送問題をどのように解けるか?
  • RQ2低ランク近似と幾何学に配慮した計算は、OT のメモリと時間計算量をどの程度削減できるか?
  • RQ3微分可能な OT は、バーチャル・バーチャルや分布間の写像といった構造的表現を学習するために利用可能か?
  • RQ4最適輸送計画を通じた暗黙的微分は、微分可能な機械学習パイプラインにおける学習安定性をどのように向上させるか?
  • RQ5グロモフ・ワッサーシュタインやソフト・ソーティングといった高度な OT 変種は、完全な微分可能性とスケーラビリティを備えて効率的に実装可能か?

主な発見

  • OTT-JAX は JAX の自動微分を完全にサポートし、輸送計画を介したバックプロパゲーションを可能にし、微分可能な最適輸送を実現しています。
  • 低ランクシンコールン・ソルバーは、特に大規模な問題において顕著なメモリおよび計算コストの削減を達成しており、精度を損なわずに行っています。
  • このツールボックスは、バーチャル・バーチャル、グロモフ・ワッサーシュタイン距離、およびソフト・ソートド配列のエンド・ツー・エンドの微分可能な計算をサポートしています。
  • 幾何クラスにより、コスト行列の暗黙的計算(例:ポイントクラウドやグリッド)が可能になり、明示的な格納を回避し、グリッド上では O(dn^{d+1}) の計算量を実現しています。
  • デロンとデゾルヌ(2020)による異なる近似法を用いて、ガウス・ミックスチャージ間のワッサーシュタインに類似した距離の効率的計算が可能になっています。
  • このツールボックスはプロダクション環境で使用可能であり、医用画像データにおける等価バーチャル計算や、らせんやスイスロールのような複雑な多様体間の形状マッチングといった高度な応用で活用されています。

より良い研究を、今すぐ始めましょう

論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。

クレジットカード登録不要

このレビューはAIが作成し、人間の編集者が確認しました。