A training run can be broken in ways that never crash it. A GPU can pass every driver check and still compute slightly wrong answers; a single slow host can drag a thousand GPUs down to its pace; the loss can blow up at step 40,000 and keep training on garbage for hours before anyone looks at the dashboard; a collective can hang so that every rank sits at 100% utilisation doing nothing. Health checks turn each silent state into a machine decision within seconds or minutes.
This article builds a health-check system for a distributed training job from first principles: what to verify before the job starts, what to measure inside the training loop, how to detect loss spikes without drowning in false alarms, how to catch hangs and stragglers, and how to respond so that every rank agrees on the action. The hardware side (ECC, Xid codes, silent data corruption, quarantine) is covered in GPU hardware faults; here the focus is on the training job and the code you control.
Three layers of health
It helps to separate three questions, because each has different signals, different latency and a different correct response.
- Is the hardware fit to start? Answered per node before scheduling and whenever a node rejoins. A wrong answer costs a whole restart, so spending minutes is worth it.
- Are the numerics sane? Answered every step from loss, gradient norm and nonfinite values. These checks must cost almost nothing.
- Is the job making progress? Answered by watchdogs and heartbeats that sit outside the training math. A hung collective produces no bad numbers, only silence.
A detector that only posts to a chat channel is a dashboard, not a health check. The goal is a closed loop: detection, a decision every rank applies identically, and a record.
Preflight: admit only fit nodes
Preflight answers whether a node should be allowed into the job at all. Three checks catch most bad nodes before they cost anything.
Vendor diagnostics. NVIDIA's dcgmi diag runs suites selected with -r: level 1 (quick) takes seconds, level 2 (medium) about two minutes, level 3 (long) about fifteen minutes and level 4 (xlong) longer still. Level 2 on every admission and level 3 after any repair is a reasonable default; the DCGM article covers the tool itself.
Communication bandwidth. The nccl-tests binaries measure collectives directly. Run all_reduce_perf inside each node and then across a pair of nodes, and compare the reported bus bandwidth against a baseline measured on known-good nodes of the same type in your own fleet; topology makes datasheet numbers useless here.
A numerical canary. Run a fixed matmul with a fixed seed on every GPU and compare a checksum across GPUs of the same model. Identical hardware running identical kernels on identical inputs should agree bit for bit; a GPU that disagrees is a candidate for silent corruption even if it passes everything else.
#!/usr/bin/env bash
# preflight.sh: exit nonzero if this node must not join the job
set -euo pipefail
dcgmi diag -r 2 || exit 10
./build/all_reduce_perf -b 8M -e 1G -f 2 -g 8 > /tmp/nccl.txt || exit 11
# check_busbw.py and canary.py are your own scripts, not part of nccl-tests
python3 check_busbw.py /tmp/nccl.txt --baseline baselines/h100_8x.json --min-ratio 0.9 || exit 12
python3 canary.py --size 8192 --seed 1234 --expect canary/h100_bf16.sha256 || exit 13# canary.py (core): same inputs, same kernel, compare bytes across GPUs
import hashlib, torch
def canary(device, n=8192, seed=1234):
g = torch.Generator(device="cpu").manual_seed(seed)
a = torch.randn(n, n, generator=g).to(device, torch.bfloat16)
b = torch.randn(n, n, generator=g).to(device, torch.bfloat16)
out = a @ b
return hashlib.sha256(out.float().cpu().numpy().tobytes()).hexdigest()
digests = {i: canary(f"cuda:{i}") for i in range(torch.cuda.device_count())}
assert len(set(digests.values())) == 1, f"GPUs disagree: {digests}"The canary compares GPUs on the same node and against a stored digest for that GPU model, driver and library version. Change any of those and the digest must be regenerated, otherwise every node fails at once after an upgrade.
In-loop numerics without slowing the step
Inside the loop, the cheapest signals are also the most informative: the loss, the global gradient norm, and whether anything is nonfinite. torch.nn.utils.clip_grad_norm_ already computes the total norm before clipping and returns it, so the gradient norm is free if you clip.
The trap is cost. Calling .item() on a GPU tensor forces the host to wait for the device, which serialises the pipeline. Keep the statistics on the GPU, combine them into one small tensor, and synchronise once per step; anything that is not needed for the skip decision can wait and be read every few hundred steps.
def train_step(model, batch, opt, health, step):
loss = model(**batch).loss
loss.backward()
gnorm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# one fused flag, still on the GPU: 1.0 if anything is nonfinite
bad = (~torch.isfinite(loss)).float() + (~torch.isfinite(gnorm)).float()
stats = torch.stack([loss.detach(), gnorm.detach(), bad])
# every rank must agree, or one rank skips while the others step and the job hangs
torch.distributed.all_reduce(stats, op=torch.distributed.ReduceOp.MAX)
loss_v, gnorm_v, bad_v = stats.tolist() # one host sync per step, not three
if bad_v > 0:
opt.zero_grad(set_to_none=True) # skip this update everywhere
return health.record(step, "skip_nonfinite", loss_v, gnorm_v)
opt.step(); opt.zero_grad(set_to_none=True)
return health.observe(step, loss_v, gnorm_v)The all_reduce with MAX is the important line. Each rank sees a different micro-batch, so only some ranks may see a nonfinite loss; if those ranks skip and the others step, the replicas diverge or the next collective hangs. Reducing the flag makes the decision global (it also makes the logged loss the worst rank's, a conservative input for spike detection). fp16 loss scalers skip overflowing steps silently; count those skips, because a rising rate is an early symptom.
At lower frequency, log per-layer gradient norms and the update-to-weight norm ratio: one layer whose gradient norm grows by orders of magnitude points at the cause long before the total norm moves.
Detecting loss spikes
A loss spike is a sudden rise that the run may or may not recover from. The PaLM paper reported about 20 spikes while training its largest model despite gradient clipping, and found that a spike came from the combination of particular batches with a particular parameter state: replaying the same batches from an earlier checkpoint did not spike. That is why the fix is to roll back and skip data, not to replay.
A fixed threshold does not work because the loss falls over the run. Compare each value to a trailing window instead, using the median and the median absolute deviation (MAD) so that the spike itself cannot inflate the baseline:
from collections import deque
import math, statistics
class SpikeDetector:
def __init__(self, window=200, z_warn=6.0, z_halt=12.0, patience=3):
self.hist = deque(maxlen=window)
self.z_warn, self.z_halt, self.patience = z_warn, z_halt, patience
self.strikes = 0
def update(self, step, loss):
if not math.isfinite(loss):
return "halt", float("inf")
if len(self.hist) < 50: # warm-up: collect a baseline
self.hist.append(loss)
return "ok", 0.0
med = statistics.median(self.hist)
mad = statistics.median(abs(x - med) for x in self.hist) or 1e-8
z = (loss - med) / (1.4826 * mad) # 1.4826 * MAD ~ sigma for normal noise
if z > self.z_warn:
self.strikes += 1
halt = z > self.z_halt and self.strikes >= self.patience
verdict = "halt" if halt else "warn"
else:
self.strikes, verdict = 0, "ok"
if verdict == "ok":
self.hist.append(loss) # spikes never enter the baseline
return verdict, zPatience matters: a single noisy batch should warn, not halt; several consecutive steps above threshold indicate real divergence. Run the detector on rank 0 after the all-reduce, then broadcast the verdict so every rank acts on the same answer.
Liveness: hangs, heartbeats and stragglers
Hangs are the most expensive failure because nothing reports them. A rank waiting in a collective for a peer that has died or diverged looks busy to every utilisation metric. PyTorch's NCCL process group has a watchdog thread for exactly this, configured by environment variables:
export TORCH_NCCL_ASYNC_ERROR_HANDLING=3 # default; error handling mode
export TORCH_NCCL_ENABLE_MONITORING=1 # abort if the watchdog itself gets stuck
export TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=600 # how long the watchdog may be silent
export TORCH_NCCL_TRACE_BUFFER_SIZE=2000 # flight recorder: ring buffer of collective events
export TORCH_NCCL_DUMP_ON_TIMEOUT=1 # dump the buffer on timeout (needs size > 0)
export TORCH_NCCL_DEBUG_INFO_TEMP_FILE=/scratch/nccl_trace_rank_Set the collective timeout on init_process_group deliberately: longer than your slowest legitimate operation, such as a checkpoint save. With the flight recorder on, a timeout leaves a per-rank record of the last collectives each rank entered; the rank that never entered the one everyone else is waiting on is your suspect.
Add an external heartbeat too. Each rank touches a file or posts a timestamp after every completed step; a supervisor outside the job kills and restarts it if the newest step is older than a few times the median step time. The same per-rank step times expose stragglers: gather them every N steps and flag any rank whose compute time (excluding time waiting in collectives) sits well above the median for several windows. Health scoring covers how to turn these signals into a node score without the average hiding a single bad host.
Responding: map each signal to an action
| Signal | Likely cause | Response |
|---|---|---|
| Nonfinite loss or gradient, isolated | overflow, bad batch | skip the update on all ranks; count skips |
| Skip rate rising | instability building | alert; consider lower LR or rollback |
| Spike: halt verdict | data plus state interaction | roll back about 100 steps, skip the data window |
| Repeated spikes after rollback | optimizer or LR too aggressive | lower LR or adjust epsilon, then resume |
| Watchdog timeout | dead or diverged rank | dump trace, abort, restart without the suspect node |
| Persistent straggler | thermal, PCIe, failing part | cordon node, run level 3 diag, restart |
| Canary mismatch | possible silent corruption | exclude node; audit recent checkpoints |
Two rules make the responses safe. First, decisions are global: one rank decides and broadcasts, or all ranks reduce a flag. Second, every automatic action writes a record with the step, ranks, data loader position and checkpoint involved. Without that record a rollback can silently skip the wrong data, and nobody can reproduce what the run actually saw. The skip itself must be implemented in the data loader state, so a resumed loader starts after the skipped window rather than at the beginning of it.
Worked example: a spike and a rollback
Take a synthetic run whose loss decays from about 3.8 towards 2.6 with Gaussian noise of 0.02, and inject a spike between steps 1000 and 1011 that peaks at +0.9 at step 1004. Running the detector above with its default settings gives:
| Step | Verdict | Robust z | Loss |
|---|---|---|---|
| 1-999 | ok (no warnings) | below 6 | 3.80 falling to 2.70 |
| 1000 | warn | 17.6 | 3.252 |
| 1001 | warn | 20.2 | 3.330 |
| 1002 | halt | 23.4 | 3.424 |
| 1004 (peak) | halt | 29.1 | 3.595 |
No warnings before the spike, a warning at the first bad step, and a halt two steps later once patience ran out. The responder then follows the PaLM-style recipe:
- Broadcast the halt; all ranks stop before the optimizer step of 1002.
- Pick the newest committed checkpoint at or before step 900 (about 100 steps before the first warning).
- Load it, then advance the data loader past the batches consumed between 900 and the spike, plus a margin (PaLM skipped roughly 200 to 500 batches).
- Resume, and keep the detector armed with a fresh baseline from the restored run.
- If the spike recurs at a different step, treat it as an optimisation problem, not a data problem.
Keeping checkpoints dense enough to roll back 100 steps is a policy decision with a real storage cost; the checkpointing deep dive covers making those saves cheap.
Failure modes of the checks themselves
- Checks that slow the job. A
.item()on every tensor you log adds a host-device sync each time. Batch statistics into one tensor and sync once. - Rank-local decisions. A rank that skips a step on its own causes a hang or silent divergence. Always reduce or broadcast the verdict.
- Poisoned baselines. If spike values enter the window, the detector adapts to the spike and goes quiet. Only append values judged normal.
- Expected discontinuities. Learning-rate warm-up, a data mixture change or a context-length switch all move the loss legitimately. Reset or widen the detector at known schedule boundaries.
- Timeouts shorter than real work. A collective timeout that fires during a long checkpoint save turns a healthy run into a restart loop.
Trade-offs
Every check trades detection speed against cost and false alarms. Level 3 diagnostics on every admission catch more bad parts but keep nodes out of service for a quarter of an hour each. A tight spike threshold catches problems sooner but halts on noise; patience and robust statistics buy precision at the price of a few steps of latency. Short collective timeouts recover hangs faster but risk killing legitimate long operations. Automatic rollback saves engineer time at night but can hide a systematic problem, so cap it and page a human when the cap is hit. Outside-in checks of the serving side are a separate discipline, covered in synthetic monitoring.
What to do next
- Write a preflight script with DCGM level 2, an nccl-tests bandwidth check against your own baseline, and a matmul canary; refuse to schedule nodes that fail.
- Make nonfinite detection global with a single all-reduced flag, and count skipped steps.
- Add a median/MAD spike detector with patience, and test it on a replay of a past run.
- Enable the NCCL watchdog and flight recorder, and set a collective timeout longer than your slowest checkpoint save.
- Add an external per-step heartbeat and a straggler report from per-rank compute time.
- Implement rollback-and-skip in the data loader state and log every automatic action.
- Cap automatic rollbacks per day and page a human when the cap is reached.