📺
YouTube · Google Cloud Tech
89 YouTube-Aufrufe
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
Vollständiger Original-Bericht
Ausführliche Details, Code-Beispiele & Hersteller-Stellungnahme auf youtube.com.