You have a training script that runs on one 8-GPU box and you need it on four, or sixteen. The model no longer fits comfortably in one GPU's memory, the input pipeline already struggles to keep eight GPUs busy, and the cluster loses a node every few days. Ray Train is one answer: it starts one worker process per GPU across the cluster, wires them into a PyTorch process group, streams data to each rank from a Ray Data pipeline, and restarts the whole group from a checkpoint when something dies.

This page is the distributed-training half of the story. The basics of Ray on GPUs, including tasks, actors, placement groups and a minimal TorchTrainer loop, are covered in Ray for GPU workloads. Here we go further: FSDP inside the worker group, data ingest, sharded checkpoints, elastic worker counts, and the arithmetic that tells you whether a 32-GPU run will be compute bound or network bound.

How a Ray Train job is put together

Think of a Ray Train job as three layers. The driver is the Python process that calls trainer.fit(). It hands the job to a controller, which owns the lifecycle: it requests resources, starts the worker group, watches it, and decides what to do on failure. The worker group is a set of Ray actors, one per GPU by default, each running your train_loop_per_worker function.

Before your function runs, Ray Train does the setup you would otherwise do by hand with torchrun: it picks a master address and port, sets rank and world size, calls torch.distributed.init_process_group with NCCL on GPUs, and restricts each worker's CUDA_VISIBLE_DEVICES to its assigned device. That is why ray.train.torch.get_device() inside the loop returns the right GPU without you computing a local rank.

Once the group is running, Ray is out of the hot path: NCCL moves gradients and shards. Ray carries only three things: data batches from Ray Data to each rank, metrics and checkpoint references reported through ray.train.report, and control messages. So when training is slow, suspect NCCL, the input pipeline or checkpoint I/O before Ray scheduling.

Ray Train on a GPU cluster: one controller, one worker group, three data pathsDrivertrainer.fit()Train controllergroup lifecycle, retriesreadsShared storageS3 / GCS / NFSstart groupNode A: 8 GPUsNode B: 8 GPUsrank 0FSDPrank 8FSDPrank 1FSDPrank 9FSDPrank 2FSDPrank 10FSDPrank 3FSDPrank 11FSDP... ranks 4-7... ranks 12-15NCCL all-gather / reduce-scatter: NVLink inside a node, RDMA between nodesRay Data streaming splitCPU decode feeds each rank's shardbatchesSharded checkpointevery rank writes its own filesuploadGradients never touch the Ray object store; only data batches and checkpoint files pass through Ray.
The controller starts one worker per GPU. Collectives run over NCCL; Ray Data feeds batches; each rank writes its own checkpoint shard to shared storage.

Placement, process groups and NCCL

Everything about placement lives in ScalingConfig. The fields that matter on GPU clusters are num_workers, use_gpu, resources_per_worker (for example one GPU and eight CPUs per worker, so the data loader has cores), placement_strategy and accelerator_type when a cluster mixes GPU models. The default strategy packs workers onto as few nodes as possible, which is what you want: ranks on the same node talk over NVLink at hundreds of gigabytes per second, while ranks on different nodes go through the NIC at a fraction of that.

Process-group settings live in TorchConfig, passed to the trainer as torch_config. It has three fields: backend (NCCL on GPU when left unset), init_method and timeout_s, which defaults to 1,800 seconds. That timeout is how long a collective may wait before the process group fails. A hung all-reduce on a dead peer then idles the whole cluster for half an hour. Lower it to a few minutes, above your slowest legitimate collective, usually the first step or a checkpoint barrier.

NCCL itself is configured with environment variables on the workers, set through the Ray runtime environment. Turning on NCCL_DEBUG=INFO for the first run of any new cluster shape is cheap insurance: the log shows which transport each ring uses, and a ring that silently fell back from RDMA to TCP sockets is the single most common reason a multi-node job runs at a third of the expected speed. The collectives themselves are explained in NCCL collectives.

Code: FSDP and sharded checkpoints inside Ray Train

The loop below trains a decoder model with FSDP across all workers, reads its data from a Ray Dataset shard, and writes a sharded checkpoint with PyTorch Distributed Checkpoint. It uses the classic FullyShardedDataParallel wrapper; the newer fully_shard API works the same way. Model construction and loss are placeholders.

import os, tempfile, functools, torch
import torch.distributed.checkpoint as dcp
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, MixedPrecision
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
import ray, ray.train
from ray.train import ScalingConfig, RunConfig, FailureConfig, CheckpointConfig, Checkpoint
from ray.train.torch import TorchTrainer, TorchConfig, get_device
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict

def save_sharded(model, opt, step, d):          # each rank writes its own shard files
    msd, osd = get_state_dict(model, opt)
    dcp.save({"model": msd, "optim": osd, "step": torch.tensor(step)}, checkpoint_id=d)

def load_sharded(model, opt, d):                 # loads into this rank's shards in place
    msd, osd = get_state_dict(model, opt)
    sd = {"model": msd, "optim": osd, "step": torch.tensor(0)}
    dcp.load(sd, checkpoint_id=d)
    set_state_dict(model, opt, model_state_dict=sd["model"], optim_state_dict=sd["optim"])
    return int(sd["step"])

def train_loop(cfg):
    device = get_device()
    model = build_model()                      # your nn.Module, on CPU or meta device
    model = FSDP(model, device_id=device,
                 auto_wrap_policy=functools.partial(transformer_auto_wrap_policy,
                                                    transformer_layer_cls={Block}),
                 mixed_precision=MixedPrecision(param_dtype=torch.bfloat16,
                                                reduce_dtype=torch.bfloat16))
    opt = torch.optim.AdamW(model.parameters(), lr=cfg["lr"])
    step = 0

    ckpt = ray.train.get_checkpoint()          # None on a fresh start
    if ckpt:
        with ckpt.as_directory() as d:
            step = load_sharded(model, opt, d)

    shard = ray.train.get_dataset_shard("train")
    while step < cfg["max_steps"]:
        for batch in shard.iter_torch_batches(batch_size=cfg["micro_bs"], device=device):
            loss = model(batch["input_ids"], labels=batch["labels"]).loss
            loss.backward()
            model.clip_grad_norm_(1.0)
            opt.step(); opt.zero_grad(set_to_none=True)
            step += 1
            if step % cfg["ckpt_every"] == 0:
                with tempfile.TemporaryDirectory() as d:
                    save_sharded(model, opt, step, d)
                    ray.train.report({"loss": loss.item(), "step": step},
                                     checkpoint=Checkpoint.from_directory(d))
            if step >= cfg["max_steps"]:
                break

trainer = TorchTrainer(
    train_loop,
    train_loop_config={"lr": 2e-4, "micro_bs": 8, "ckpt_every": 500, "max_steps": 20_000},
    scaling_config=ScalingConfig(num_workers=32, use_gpu=True,
                                 resources_per_worker={"GPU": 1, "CPU": 8}),
    torch_config=TorchConfig(timeout_s=600),
    datasets={"train": ray.data.read_parquet("s3://bucket/tokens/")},
    run_config=RunConfig(storage_path="s3://bucket/runs", name="lm-7b",
                         failure_config=FailureConfig(max_failures=5),
                         checkpoint_config=CheckpointConfig(num_to_keep=3)),
)
result = trainer.fit()

Three choices in that code are deliberate. First, every rank calls ray.train.report with its own checkpoint directory. Ray Train uploads each worker's files independently into the same checkpoint location in storage, and files with the same name overwrite one another, so the save must use rank-unique names. Distributed Checkpoint does that for you: each rank writes its own shard files plus shared metadata. Second, nothing gathers the full state dict to rank 0. For a 7B model with AdamW the full state is over 100 GB, which would not fit in one GPU and would take minutes to gather. Third, num_to_keep stops the bucket from growing by one full state per checkpoint. The trade-offs between sharded and consolidated formats are covered in GPU checkpointing, and the FSDP wrapping options in PyTorch FSDP.

Getting data to every rank

When a dataset is passed through datasets, Ray Data splits it into one stream per worker and runs the read, decode and transform stages on CPU workers across the cluster, pipelined ahead of the GPUs. get_dataset_shard returns that worker's stream; iter_torch_batches collates and moves each batch to the device. Preprocessing scales with cluster CPUs rather than with the cores on each GPU node, and you do not need a DistributedSampler, because sharding already happened.

Size the pipeline before blaming the model. In the worked example below each rank processes 32,768 tokens per 3.4-second step, so 32 ranks consume about 310,000 tokens per second: a few cores of tokenisation at, say, 100,000 tokens per second per core. A model ten times smaller at the same batch steps about ten times faster and needs over 30 cores for tokenising alone, so measure your tokenizer's real rate. If GPU utilisation shows a saw-tooth with gaps between steps, the shard is running dry: add CPU nodes, pre-tokenise, or raise the prefetch depth.

Worked example: will 32 GPUs be compute bound?

Worked example: a 7-billion-parameter decoder, trained in bf16 with fp32 AdamW state, on four nodes of eight H100 SXM GPUs, each node with eight 400 Gb/s NICs.

QuantityArithmeticResult
Training state16 bytes per parameter (bf16 weights and grads, fp32 master copy, two Adam moments) x 7e9about 112 GB
State per GPU under full sharding112 GB / 323.5 GB plus activations
FSDP traffic per GPU per stepabout 3 x 2 bytes x 7e9 (forward gather, backward gather, reduce-scatter)about 42 GB
Time on the NIC42 GB / 50 GB/s per GPUabout 0.84 s
Compute per GPU per step6 x 7e9 x 32,768 tokensabout 1.4e15 FLOP
Time at 40% of 989 TFLOP/s dense bf161.4e15 / 4e14about 3.4 s

Communication at 0.84 seconds against 3.4 seconds of compute means FSDP's prefetching can hide almost all of it: the job should be compute bound. Halve the micro-batch and compute drops to 1.7 seconds while traffic stays at 0.84, so overlap gets tight. Run the same job over a single 100 Gb/s Ethernet port per node and traffic time rises roughly thirty-fold, turning a compute-bound job into a network-bound one. This is why the arithmetic comes before the cluster order, and why hybrid sharding (full sharding inside a node, replication across nodes) is worth trying when inter-node bandwidth is the constraint.

Checkpoints get the same treatment. Writing 112 GB at an aggregate 2 GB/s to object storage takes roughly a minute. If the 32-GPU cluster's mean time between failures is eight hours (28,800 seconds), Young's approximation for the interval that minimises lost work is the square root of 2 x 60 x 28,800, about 1,860 seconds. Checkpointing every 30 minutes, not every 5 or every 3 hours, is the right order of magnitude; at 3.4 seconds per step, that is about every 500 steps, which is where the example's ckpt_every came from.

Failures, restarts and elastic worker counts

What a failure costs: the whole group restarts from the last reported checkpointTrainingsteps advanceckptTrainingwork since ckptGPU XidGroup torn downall ranks stopNew groupsame or fewer nodesload ckptResumelost = work since ckptCost per failure = restart time + checkpoint load + half an interval of work,so the checkpoint interval is a dial you tune against failure rate.
Recovery is group-wide: every rank stops, a new group starts, and all ranks load the latest reported checkpoint.

When any worker dies, from a GPU Xid error, a host OOM or a lost node, Ray Train shuts down every worker in the group, acquires resources again, starts a fresh group, and each rank loads the latest reported checkpoint through ray.train.get_checkpoint(). This happens up to max_failures times; the default of zero means no recovery at all, and minus one retries forever. Group-wide restart is not a design flaw: a synchronous data-parallel job cannot continue with a hole in its collectives, so replacing one rank means rebuilding the process group anyway.

Recent Ray Train releases also support elastic worker counts: pass num_workers=(8, 32) and Ray requests the maximum, starts as soon as the minimum is available, and resizes the group, again by restarting from a checkpoint, when the autoscaler adds nodes. The docs pair it with max_failures and an autoscaler configured to match. Check that your Ray version supports it, since it belongs to the newer Train implementation. Elasticity changes the global batch size unless you compensate: either keep the global batch fixed by raising gradient accumulation on smaller groups, or rescale the learning rate, and log the world size with every metric so loss curves can be interpreted later.

Failure modes

SymptomLikely causeWhat to do
Step time 3x the estimate on multi-node onlyNCCL fell back to TCP sockets; RDMA devices not visible in the containerRun once with NCCL_DEBUG=INFO, check the transport line, fix device plugins or network env vars
Run hangs, then fails 30 minutes laterA peer died mid-collective and timeout_s is at its defaultLower timeout_s in TorchConfig; alert on step-time stalls
Restart loads stale or mixed weightsRank files with identical names overwrote each other in storageUse DCP or a rank suffix in every file name
Restart succeeds, loss jumpsData stream restarted from the beginning or optimizer state not restoredCheckpoint optimizer and step; skip already-seen data using the step count
Host OOM on the GPU nodeFull state dict gathered on rank 0 for savingSave sharded; consolidate offline if a single file is needed
GPUs idle between stepsRay Data shard drained; too few CPUs for decodeMore CPU nodes, pre-tokenise, increase resources_per_worker CPU

Trade-offs: Ray Train or plain torchrun

Ray Train is not the only way to run this job. Plain torchrun under Slurm or a Kubernetes training operator gives you the same NCCL data path with one less control plane. If the job is a fixed-size pretraining run on a dedicated cluster with a fast parallel filesystem, that simplicity often wins.

Ray earns its place when training is one stage among several: CPU preprocessing scaled separately, Ray Tune sweeps, reinforcement learning, or nodes shared with serving. The costs are a head node to protect and one more layer to debug through when a collective hangs. For models too large for FSDP alone you can run DeepSpeed inside the same loop, since Ray only provides the process group.

What to do next

  1. Do the table for your own model before ordering hardware: state per GPU, FSDP bytes per step, NIC time and compute time. If NIC time exceeds about half the compute time, plan for hybrid sharding or larger micro-batches.
  2. Move one job to TorchTrainer with FSDP and Distributed Checkpoint, and confirm every rank's files land in the same checkpoint directory in storage.
  3. Run the first multi-node job with NCCL_DEBUG=INFO and verify the transport is RDMA, not sockets.
  4. Set timeout_s in TorchConfig to a few minutes and max_failures to a positive number, then kill a worker node mid-run and time the recovery end to end.
  5. Compute your checkpoint interval with Young's formula from measured save time and observed failure rate, and revisit it each quarter.
  6. Feed data through Ray Data shards, watch for gaps in GPU utilisation, and scale CPU nodes until they disappear.
Key takeaway: Ray Train builds the PyTorch process group for you and then gets out of the hot path: NCCL carries gradients, Ray Data carries batches, and each rank writes its own checkpoint shard. Do the arithmetic first, since state per GPU, FSDP bytes against NIC bandwidth, and the Young checkpoint interval decide whether the run is fast and recoverable. Then set timeout_s and max_failures, save sharded, and test recovery by killing a node on purpose.