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 記憶體執行。

適用對象

需要使用 JAX 在多個主機與 GPU 上進行擴展的科學計算工作負載的研究人員與開發者。

主要亮點

  • JAX 數組的零複製通訊。
  • jax.jit 相容,實現高效率執行。
  • 支援跨 CPU 與 GPU 集群的多主機擴展。
  • 支援對某些 MPI 操作進行微分,適用於可微分程式設計。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案