Skip to content

Large distributed training#

A training run that needs more than a handful of GPUs has to survive the things that happen at that scale: a higher-priority run needs the capacity, a machine goes down mid-step, a cable is a little slower than its peers. This recipe runs PyTorch DDP as a gang across 8 machines × 8 GPUs (64 GPUs), asks Astraeus to keep it on as few switch hops as there is room for, checkpoints to a shared drive, and comes back on its own — on the same machines or different ones — after either kind of interruption.

Before you begin#

  • Eight or more 8-GPU machines in your workspace's pool, on one InfiniBand fabric. Machines with RDMA run their workers on the host network, which is what NCCL over InfiniBand needs.
  • The editor role in the workspace, an API token, curl and jq (How the recipes are written).
  • A drive both the run's machines can reach for checkpoints — a drive on each machine is enough here, since rank 0 alone writes (see Drives).

How Astraeus places it, and keeps it there#

With start: Gang, every worker is placed in one decision or none: the run either gets all 8 machines at once or stays Pending. Placed whole, a gang climbs this ladder and stops at the first rung every worker fits on:

  1. One machine · 2. one NVLink domain · 3. one leaf group (one hop on every rail) · 4. one InfiniBand fabric and one rack · 5. one InfiniBand fabric · 6–9. progressively looser RDMA and Ethernet.

topology.keep_within: fabric keeps the whole gang within one InfiniBand fabric — never split across two fabrics that cannot talk over RDMA — and patience_seconds lets it wait for a tighter rung (one leaf group) rather than take the first fit across the spines:

Waiting up to 25 min more for one leaf switch (leaf-r0-7-g0), ~368 GB/s all-reduce (could start now on InfiniBand fabric sm-0x0002c90300a1b2c3, ~313 GB/s all-reduce)

Once placed, the run's network field names what it got (leaf:<leaf group> or fabric:<id>) and the all-reduce bandwidth to expect there, so you can tell a run that is merely working from one that is working at the fabric's full speed. See How placement uses it and The expected all-reduce bandwidth.

1. Write a checkpointing training script#

Every worker runs torchrun, which starts 8 processes (one per GPU); the global rank, not the Astraeus worker's rank, is what matters for which process saves. Rank 0 of the whole run writes the checkpoint, atomically, on SIGTERM and periodically:

train_ddp.py
import os, signal, sys, threading
import torch
import torch.distributed as dist
import torch.nn as nn
import torchvision

CKPT_DIR = "/ckpt/" + os.environ.get("ASTRAEUS_RUN_NAME", "run")
CKPT = os.path.join(CKPT_DIR, "latest.pt")
stop = threading.Event()
signal.signal(signal.SIGTERM, lambda *_: stop.set())


def save(step, model, opt):
    os.makedirs(CKPT_DIR, exist_ok=True)
    tmp = CKPT + ".tmp"
    torch.save({"step": step, "model": model.state_dict(), "opt": opt.state_dict()}, tmp)
    os.replace(tmp, CKPT)  # atomic: a kill mid-save keeps the previous one


def main():
    dist.init_process_group("nccl")
    rank, world = dist.get_rank(), dist.get_world_size()
    local = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local)
    dev = torch.device("cuda", local)

    model = nn.parallel.DistributedDataParallel(torchvision.models.resnet50().to(dev), device_ids=[local])
    opt = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9)
    loss_fn = nn.CrossEntropyLoss()
    start_step = 0
    if os.path.exists(CKPT):
        state = torch.load(CKPT, map_location=dev)
        model.module.load_state_dict(state["model"])
        opt.load_state_dict(state["opt"])
        start_step = state["step"] + 1
        if rank == 0:
            print(f"resuming at step {start_step}", flush=True)

    images = torch.randn(256, 3, 224, 224, device=dev)
    labels = torch.randint(0, 1000, (256,), device=dev)
    total_steps = int(os.environ.get("STEPS", "50000"))
    for step in range(start_step, total_steps):
        loss = loss_fn(model(images), labels)
        opt.zero_grad()
        loss.backward()
        opt.step()
        if stop.is_set():
            if rank == 0:
                save(step, model, opt)
                print(f"stopped at step {step}, checkpoint saved", flush=True)
            dist.barrier()
            sys.exit(0)
        if rank == 0 and step % 200 == 0:
            save(step, model, opt)
            print(f"step {step}/{total_steps} loss {loss.item():.3f} on {world} GPUs", flush=True)
    dist.destroy_process_group()


if __name__ == "__main__":
    main()

This saves the full five points of Write a checkpoint-friendly run: a drive, not the container's filesystem; resume at start; save on SIGTERM within the stop grace, and periodically for a lost machine, which sends no signal; an atomic write; and exec so the shell forwards the signal (below).

2. Write the run specification#

train-64gpu.json
{
  "metadata": {"name": "resnet-64gpu", "labels": {"project": "vision-scale"}},
  "spec": {
    "start": "Gang",
    "on_failure": "RestartJob",
    "task_template": {
      "image": "nvcr.io/nvidia/pytorch:24.08-py3",
      "command": "bash",
      "args": ["-c", "exec torchrun --nnodes=$WORLD_SIZE --nproc-per-node=$ASTRAEUS_GPU_COUNT --node-rank=$RANK --master-addr=$MASTER_ADDR --master-port=$MASTER_PORT /app/train_ddp.py"],
      "env": {"STEPS": "50000"},
      "time_limit_seconds": 86400,
      "drives": [{"name": "checkpoints", "mount_path": "/ckpt", "mode": "ReadWrite"}],
      "requested_resources": {
        "gpu_requests": {"count": 64, "models": ["H100"], "healthy_only": true},
        "machines": 8,
        "per_gpu": {"cpu_cores": 12, "memory_bytes": 107374182400},
        "network": {"interconnect": "infiniband", "min_gbps": 400},
        "topology": {"keep_within": "fabric", "patience_seconds": 900},
        "machine_selection": {"mode": "Any"}
      },
      "configs": []
    }
  }
}
$ jq --rawfile src train_ddp.py \
    '.spec.task_template.configs = [{"mounts": ["/app/train_ddp.py"], "value": $src}]' \
    train-64gpu.json > train-64gpu.full.json
Field Why
gpu_requests.count: 64, machines: 8 64 GPUs over exactly 8 machines, one worker each; the count must divide by machines.
on_failure: RestartJob Synchronous training does not survive the loss of a member: any worker failing, or its machine going down, restarts the whole gang.
drives: [{mode: "ReadWrite"}] A drive on each machine, mounted at /ckpt, read-write, for the checkpoint.
network.interconnect: infiniband, min_gbps: 400 Only machines with 400 Gb/s InfiniBand ports; the worker gets only the ports up and at that rate in NCCL_IB_HCA.
topology.keep_within: fabric Never split across two fabrics that cannot reach each other over RDMA.
topology.patience_seconds: 900 Wait up to 15 minutes for a tighter placement (ideally one leaf group) than the one free now. Other work runs meanwhile.
healthy_only Never place on a GPU already showing signs of a coming fault.

3. Submit it#

Runs → New run → Edit as JSON, paste train-64gpu.full.json, and Start run.

Save the specification as a workspace template (console: Save as template, or the templates API), then:

$ astra astraeus run --template resnet-64gpu --name resnet-64gpu-run1
resnet-64gpu-run1 submitted to lab-a: https://console.astralyx.cloud/o/acme/w/vision/runs/lab-a/resnet-64gpu-run1
$ curl -fsS -X POST "$ASTRALYX_API/runs" -H "Authorization: Bearer $ASTRALYX_TOKEN" \
    -H 'content-type: application/json' --data-binary @train-64gpu.full.json | jq -r .metadata.name
resnet-64gpu

4. Watch it start, and check the network it got#

$ astra astraeus runs
NAME            CLUSTER  STATE     REASON
resnet-64gpu    lab-a    Starting  Starting together: 6 of 8 ready
$ curl -sS "$ASTRALYX_API/runs/resnet-64gpu" -H "Authorization: Bearer $ASTRALYX_TOKEN" | jq .network
{
  "network": "leaf:leaf-r0-7-g0",
  "words": "one leaf group",
  "estimate": {"busbw_gbps": 368.4, "words": "one leaf switch", "basis": "8 × 400 Gb/s rails, × 0.92"}
}

If it waited out the full patience and settled for crossing the spines instead, words says 2 leaves under one spine group and the estimate is lower — still correct, just a looser placement. See The expected all-reduce bandwidth.

5. Surviving preemption#

A higher-priority run needs these GPUs. Its workers get SIGTERM, the machine's stop grace to save (rank 0 writes the checkpoint), then the gang goes back to the queue:

$ astra astraeus runs
NAME          CLUSTER  STATE    REASON
resnet-64gpu  lab-a    Pending  Preempted by higher-priority work (eval-release); waiting to run again

Preemption does not spend the run's restart budget — only real failures do (below). Once room returns, the gang restarts and resumes from /ckpt/resnet-64gpu/latest.pt:

$ astra astraeus logs resnet-64gpu-0 | tail -n 3
resuming at step 18400
step 18400/50000 loss 2.104 on 64 GPUs

See Preemption.

6. Surviving a lost machine#

A machine carrying one of the eight workers goes Down — a power fault, a kernel panic. With on_failure: RestartJob, the whole gang is requeued and placed again, on whichever 8 machines fit (the same seven plus a replacement, or a different group entirely):

$ astra astraeus runs
NAME          CLUSTER  STATE    REASON
resnet-64gpu  lab-a    Pending  Gang restart: resnet-64gpu-5 lost its machine

Unlike preemption, this does spend the restart budget: restarts back off from 10 s, doubling to 5 minutes, with a budget of 10 — a run running 10 minutes earns a restart back. Training resumes from the same checkpoint once the gang starts again. The lost machine itself goes through self-healing in parallel: faults are found, the machine is kept out of new work, and repaired.

NCCL_ASYNC_ERROR_HANDLING=1, set on every GPU worker by default, makes a collective that loses a peer fail fast instead of hanging until the time limit — so the gang restarts promptly rather than sitting idle.

7. Clean up#

$ astra astraeus delete resnet-64gpu

Variations#

Let Astraeus choose the shape. Drop machines and set "gpu_requests": {"count": 64, "min_per_machine": 8}: the scheduler picks the fewest machines on the best network with room and may change its mind until the run starts. RANK, WORLD_SIZE and ASTRAEUS_GPU_COUNT always describe the shape it chose.

One NVLink domain (GB200 NVL72). "topology": {"keep_within": "nvlink_domain"} with "network": {"interconnect": "nvlink"} keeps every worker in one multi-node NVLink domain, sharing GPU memory across the machines through IMEX channel 0.

RoCE instead of InfiniBand. "network": {"interconnect": "rdma"} accepts InfiniBand or RoCE, never plain Ethernet.

Pin to a known leaf group. When you already know a leaf group's name from Topology (leaf-r0-7-g0), you can select it directly: "machine_selection": {"mode": "Any", "match_labels": {"topology.astraeus.io/leaf": "leaf-r0-7-g0"}}. Most runs do better to let keep_within: fabric and patience_seconds find the best leaf group available, since a hard pin waits only for that one.

Troubleshooting#

Symptom Cause Fix
The run restarts from step 0 after a failure The checkpoint did not save before the worker died, or start_step logic is wrong. Save periodically, not only on SIGTERM: a lost machine sends no signal.
Gang barrier timed out waiting for all members A worker was still pulling the image 15 minutes after the first reached the barrier. Pre-pull large images; use a smaller one.
The log shows NET/Socket, not NET/IB Placed without RDMA. Keep network.interconnect: infiniband so the run waits for a fabric rather than fall back.
Bus bandwidth well below the run's network.estimate A degraded or miswired link on one of the run's machines. Read Topology's findings about those machines.
torch.distributed.DistStoreError: … timed out on a non-leader worker The leader's torchrun was not up yet, or the rendezvous port is in use. Set "env": {"MASTER_PORT": "29600"}.
The run keeps restarting and eventually fails: the gang does not run without it The restart budget (10) is spent: a real, recurring failure, not preemption. Check the failing worker's machine for a hardware fault (Self-healing).