Scale JAX models to multi-GPU systems @googlecloudtech
Scale JAX models to multi-GPU systems  @googlecloudtech
Uploaded August 2026 | Updated September 2026, 2 weeks ago
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
Scale JAX models to multi-GPU systemsGenerative UI for any agent, anywhere: A2UI, AG-UI, MCP Apps, and moreBuild long-running agents with Google’s Agentic Stack | The Agent Factory5 agent patterns to masterHow Google Developer Experts vibecoded an AI racing coach with GeminiAlloyDB trial clusters: your playground for PostgreSQL innovationBuilding distributed multi-agent systemsBuild your first Android app in AI Studio in 5 minutesData agent kit: Your coding agent can now query your dataConverting SQL Server types and rules to PostgreSQLBuild a multi-agent system using ADK & MCPWhat is an Agentic Harness?
Google Cloud Tech |

Scale JAX models to multi-GPU systems

SHARE TO X SHARE TO REDDIT SHARE TO FACEBOOK WALLPAPER