mpi4jax/mpi4jax

Zero-copy MPI communication of JAX arrays, for turbo-charged HPC applications in Python :zap:

何を解決するか

JAXフレームワークのマルチホスト機能が限られている問題に対処し、ユーザーがJAXベースのシミュレーションをCPUおよびGPUクラスタ全体にスケーリングできるようにします。

動作方法

JAX配列のゼロコピー、マルチホスト通信を可能にします。この機能はJAXのジャストインタイム(jit)コンパイルに直接統合されているため、jittedコード内から通信が可能であり、GPUメモリから直接通信が行えます。

対象ユーザー

複数のホストおよびGPUを用いてJAXでスケーリングが必要な科学計算ワークロードを実行する研究者および開発者。

特徴

  • JAX配列のゼロコピー通信。
  • 高性能実行のための jax.jit との互換性。
  • CPUおよびGPUクラスタ全体でのマルチホストスケーリングをサポート。
  • 特定のMPI操作を微分可能プログラミングに適用できるようにサポート。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト