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 操作進行微分,適用於可微分程式設計。
相關
- 專案
- 專案
- 專案
- 專案
- 專案