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操作を微分可能プログラミングに適用できるようにサポート。
関連
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト