learninfra · Linux · Networking · Kubernetes · System Design · AI Infrastructure · Exam blueprints · Drills

Ray & Distributed AI Frameworks

The layer between the cluster and the model code: PyTorch distributed and NCCL, DeepSpeed, Megatron-LM and FSDP, Ray and KubeRay, JAX on TPUs, who owns what, and how a distributed launch fails.

An interactive AI Infrastructure lesson: 21 steps, about 30 minutes, on a live simulation in your browser.

The ML team wants to fine-tune a Llama 3.1 8B-shaped model on 16 GPUs. The platform team runs Kubernetes 1.35 with the NVIDIA GPU Operator on four DGX H100s. Between them sits a third layer that neither team wrote: a distributed framework that turns 16 processes on two machines into one training program.

Submit it. Kubernetes places two worker pods, ft-worker-0 on dgx-1 and ft-worker-1 on dgx-2, 8 GPUs each. Inside each pod the framework starts 8 processes, one per GPU, and wires all 16 into a group that averages gradients every step: 6.54 s per step, MFU 49%.

What you will learn

  1. The layer in between

    • Sixteen processes, one program: The cluster provides machines and GPUs; the framework turns processes on them into one program; the model code only sees tensors.
    • Who owns what: The contract between platform and ML teams is an image, a job spec and resource requests. Everything outside it belongs to exactly one team.
  2. PyTorch distributed

    • Ranks, world size, process group: A distributed PyTorch job is WORLD_SIZE identical processes; each knows its RANK, and they only talk through collectives on a process group.
    • What NCCL asks of the platform: NCCL is only as fast as the slowest hop it can see from inside the container: run nccl-tests with the same image and pod spec the jobs will use.
    • Drill: launch 16 ranks with torchrun
  3. DeepSpeed, Megatron, FSDP, JAX

    • Making 70B fit on 32 GPUs: Sharding frameworks are memory tools first: ZeRO/FSDP split the model state, checkpointing trades compute for activations, tensor and pipeline parallelism split the layers themselves.
    • DeepSpeed: ZeRO in a config file: DeepSpeed is ZeRO sharding plus a launcher, driven by one JSON config: the stage decides what every GPU stops holding.
    • Megatron-LM: 3D parallelism: Megatron-LM splits a model three ways: tensor inside a server, pipeline across servers, data across copies. tp × pp × dp must equal the number of GPUs.
    • PyTorch FSDP: sharding built in: FSDP is ZeRO-3 in PyTorch: shard per block, gather just in time, free right after. Memory falls with the number of ranks; traffic rises by half.
    • The same ideas in JAX on TPUs: JAX moves the parallelism decisions into the compiler: you declare a device mesh and array shardings, XLA writes the collectives.
  4. Ray and KubeRay

    • Ray: a distributed runtime: Ray is a cluster-wide Python runtime: a head that coordinates, workers that execute, and an object store that moves results between them.
    • Tasks, actors and Ray Train: A task is a remote function call, an actor is a remote object with state; Ray Train is a group of actors running an ordinary torch.distributed job.
    • Ray Serve and Ray Data: Ray Data, Train and Serve are libraries on one runtime; Ray schedules inside the resources the cluster scheduler has already handed it.
    • KubeRay: Ray on Kubernetes: KubeRay turns a RayCluster resource into pods: Ray decides it needs more workers, Kubernetes decides whether they get GPUs.
  5. A failed launch

    • A job that half-starts: A distributed job needs all its ranks or none. A scheduler that places pods one by one lets it hold half its GPUs and do nothing: use gang scheduling.
    • Break it: a rank nobody can reach: Running pods are not a running job. A rank that cannot reach its peers turns into a hang, and the orchestrator reports everything healthy.
    • Find the rank, fix the path: Diagnose a stuck launch from the outside in: pods, NCCL logs, nccl-tests on the suspect nodes, then the NIC. Fix the path, then restart; a communicator does not recover.
    • Break it: an NCCL setting from another cluster: A wrong NCCL setting usually fails loudly at init; a missing network path hangs quietly later. Both look like the ML team's bug and both are the platform's to prevent.
    • Drill: make NCCL explain itself
  6. Recap & playground

    • Cheat sheet
    • Playground