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 操作进行微分,适用于可微编程。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目