Two implementations of ZeRO-1 optimizer sharding in JAX
-
Updated
Jun 11, 2023 - Python
Two implementations of ZeRO-1 optimizer sharding in JAX
Tensor Parallelism with JAX + Shard Map
Comprehensive JAX implementation of neural networks and scientific computing. Features distributed training, physics-informed networks, custom autodiff, and advanced optimization. Production-ready code with numerical stability, multi-device parallelism, and research-grade implementations.
To associate your repository with the pjit topic, visit your repo's landing page and select "manage topics."