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