Scale JAX models to multi-GPU systems
Join the Google Cloud & NVIDIA community → https://g.dev/cloud/google-nvidia-community
Scaling deep learning models across multiple GPUs used to mean rewriting hundreds of lines of complex device communication code.
This is part 3 of JAX on NVIDIA GPUs Crash Course.
Watch along and learn about JAX's modern, compiler-driven sharding model, demonstrating how to distribute workloads automatically across physical device meshes.
* *Master sharding concepts:* Understand how Mesh, PartitionSpec, and NamedSharding declare array layouts across multiple devices. * Implement automatic scaling: Write clean training code that allows the compiler to automatically manage multi-GPU gradient synchronization.
* *Incorporate Flax NNX & Orbax:* See how to manage state and serialize model checkpoints in a distributed training run.
Watch more JAX on NVIDIA GPUs Crash Course → https://g.dev/cloud/jax-nvidia-gpu
? Subscribe to Google Cloud Tech → https://goo.gle/GoogleCloudTech
Speakers: Ivan Nardini, Ekaterina Sirazitdinova
Products Mentioned: Google Cloud, JAX
Google Cloud Tech
Helping you build what's next with secure infrastructure, developer tools, APIs, data analytics and machine learning....