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,
curlandjq(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:
- 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:
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#
{
"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:
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#
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). |