Cloud Architecture

Fault-Tolerant Distributed AI Training: TorchElastic, TorchFT, Distributed Checkpointing, and Keeping Your GPU Job Running When Hardware Refuses to Cooperate

A practical guide to fault-tolerant distributed GPU training: TorchElastic, TorchFT replica groups, PyTorch DCP, async checkpointing, and the data pipeline checkpoint that most teams forget.

Diagram showing distributed GPU training nodes with fault tolerance layers including checkpointing, elastic restart, and TorchFT replica groups

I have been building distributed computing systems for twenty years, and nothing humbles you quite like watching a 96-hour training job die at the 91-hour mark because one H100 in rack 14 decided it was done. GPU hardware at scale is not reliable in the way that, say, a Postgres replica is reliable. The physics are against you: dense packing, high thermal loads, exotic memory technologies, and interconnects running at the edge of what copper and silicon can sustain. When you train a frontier-scale model, hardware failures are not edge cases you plan for in a disaster recovery document. They are scheduled interruptions you build your training loop around.

I want to talk about how to do that properly. The tools have matured significantly in 2026. PyTorch now gives you three distinct recovery strategies with different trade-offs: checkpoint-restart with distributed asynchronous saves, elastic restart through TorchElastic, and checkpoint-less recovery through TorchFT. Each solves a different failure profile. Picking the wrong one for your training regime will cost you either engineering complexity you do not need or GPU-hours you cannot afford to lose.

Why GPU Hardware Fails More Than You Expect

Before diving into recovery strategies, it helps to understand why you need them in the first place. The failure modes at GPU training scale fall into a few categories.

XID errors from the GPU driver indicate hardware fault events: memory errors, NVLink issues, PCIe problems, power delivery anomalies. Some XIDs are recoverable and correctable; others signal a hard fault that will take the process down. At the density of modern H100/H200 clusters, the aggregate MTBF of any given node is short enough that in a large training run you should expect multiple node-level failures over the course of the job.

Network fabric issues are the other major source. InfiniBand and RoCE fabrics at scale generate transient errors: packet loss events, switch resets, link flaps. Collective communications (AllReduce, AllGather, ReduceScatter) are deeply sensitive to any process that falls behind or disappears. A single straggler in a 512-GPU all-reduce will stall every other rank waiting for its contribution. The NCCL watchdog will eventually time out and kill the process group, which takes your entire training job down with it.

There is also plain software: out-of-memory kills, driver version mismatches after a hot patch, CUDA kernel panics on malformed input batches. A real production training cluster deals with all of these.

The point is not to catalogue failures. The point is that at scale, the failure rate is high enough that your recovery strategy is part of your training architecture, not an afterthought.

Diagram of GPU cluster failure taxonomy showing XID errors, network fabric issues, and OOM kills as the three primary failure modes in large distributed training jobs

The Baseline: torch.save Is Not a Recovery Strategy

The most common approach I see in teams starting out is periodic saves using torch.save. Every N steps, rank 0 calls torch.save on the model state dict and writes it to shared storage. When the job dies, you relaunch and load the checkpoint.

This works, with two significant problems.

First, torch.save on a large model has rank 0 collect the full state dict, which requires peak memory roughly equal to the model size on a single GPU before writing. On a 70B parameter model, that is the entire thing. This limits checkpoint frequency because each save is expensive both in memory and in write time to shared storage.

Second, when you save infrequently to keep costs manageable, every failure costs you all the training steps since the last checkpoint. With a 6-hour save interval and failures every few hours at scale, your effective goodput, the fraction of compute time that actually advances training, can drop dramatically.

The good news is that PyTorch’s Distributed Checkpoint API (DCP), which has been stable and production-worthy for some time now, eliminates the first problem almost entirely.

PyTorch Distributed Checkpoint (DCP): Parallel Sharded Saves

DCP changes the model for saving state: instead of collecting everything on rank 0, each rank writes its own shard in parallel. The API is built around torch.distributed.checkpoint.save and torch.distributed.checkpoint.load, with a storage backend (local filesystem, object storage, or a vendor-specific backend like Databricks’ UCVolumeWriter).

import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict

# Saving
state = {"model": model, "optimizer": optimizer}
model_state, optimizer_state = get_state_dict(model, optimizer)
dcp.save(
    {"model": model_state, "optimizer": optimizer_state},
    storage_writer=dcp.FileSystemWriter(checkpoint_path),
)

# Loading (can resume on a different GPU count)
model_state, optimizer_state = get_state_dict(model, optimizer)
dcp.load(
    {"model": model_state, "optimizer": optimizer_state},
    storage_reader=dcp.FileSystemReader(checkpoint_path),
)
set_state_dict(model, optimizer, model_state=model_state, optim_state=optimizer_state)

What makes DCP valuable in production is that the metadata it writes allows resharding: you can save a checkpoint from a 256-GPU run and resume it on 128 or 512 GPUs. That flexibility matters when you are restarting a failed job onto a partially reduced cluster or when you need to continue training on different hardware after a procurement change. The parallel file systems for AI training article covers the storage layer that makes this fast in practice.

Async Checkpointing: Making Frequent Saves Nearly Free

DCP parallelizes writes across ranks. Async checkpointing takes it further by staging the checkpoint to a fast local buffer while training continues. The GPU writes do not block the training loop for more than the time it takes to copy tensors to pinned host memory; the actual I/O to persistent storage happens in the background.

The Databricks AI Runtime documentation describes this pattern explicitly. Their UCVolumeWriter stages through a fast local NVMe-backed directory before uploading to Unity Catalog volumes. The key discipline: make sure your async writer marks completion only after all data has reached its destination, not just after the local stage. I have seen teams lose checkpoints because their writer reported success after the local copy and the background upload was interrupted by the same failure that killed the training job.

With async checkpointing, saving every 100 steps rather than every 1000 steps becomes operationally reasonable. The mean time to recovery from a failure drops proportionally. This is the core of what Databricks calls “goodput”: the fraction of your GPU-hours that contribute to forward training progress. Frequent async saves move the recovery cost from hours of lost computation to minutes.

One more thing that teams consistently miss: checkpoint your data loader state alongside your model. If you restart a job from step 5000 but your data loader starts from the beginning of the dataset, you are silently re-exposing your model to the same data it has already seen. For large datasets with complex sampling strategies, this distorts the effective training distribution without producing any error or warning. Save your data pipeline’s position explicitly.

TorchElastic: Elastic World-Size Recovery

The previous strategies cover what to save and how to save it cheaply. TorchElastic (launched via torchrun) addresses a different problem: what happens when the number of available GPUs changes mid-run.

Without elastic restart, a training job defines its world size at launch and that is final. If you lose a node, the job dies and you relaunch. With TorchElastic, the job defines a range for world size (--min-nodes, --max-nodes) and the Elastic Agent on each node coordinates through a Rendezvous backend to reform the process group when membership changes.

The model requires your training script to be idempotent with respect to initialization: every time the process group reforms, dist.init_process_group runs again, the model is re-sharded across the new world size, and training resumes from the last checkpoint. The agent handles rendezvous (which nodes are in the group), and your script handles state restoration.

torchrun \
  --nnodes 8:12 \
  --nproc-per-node 8 \
  --rdzv-backend etcd \
  --rdzv-endpoint etcd-server:2379 \
  --rdzv-id job-42 \
  train.py

The 8:12 range tells the launcher to proceed with anywhere from 8 to 12 nodes. If you lose a node, the remaining nodes reform and continue. If a replacement node comes online, it can join the running job.

The important operational constraint: torchrun supports only homogeneous local world sizes. Every node must run the same number of workers per role. You cannot have some nodes with 8 GPUs and others with 4. In cloud environments where you are using uniform instance types this is usually fine; in heterogeneous on-premises clusters it can be limiting.

TorchElastic pairs well with Kubernetes and Karpenter for handling the node lifecycle side. When Karpenter provisions a replacement spot instance, TorchElastic can absorb it into the running job without a full restart.

Architecture diagram showing TorchElastic Elastic Agent and Rendezvous backend coordinating multiple nodes that dynamically join and leave a training job while checkpointing handles state persistence

TorchFT: Checkpoint-Less Recovery with Replica Groups

TorchFT is a newer and architecturally distinct approach. Where TorchElastic treats the entire training run as a single elastic process group, TorchFT structures the job as multiple parallel replica groups that each hold a complete copy of the model (using DDP or HSDP/FSDP with DDP across groups).

The architecture has two core components. A global Lighthouse server tracks the health of all replica groups through heartbeats. Each replica group runs a Manager that handles coordination with the Lighthouse. When a replica group fails, the Lighthouse removes it from the active quorum. Training continues with the remaining groups. When a replacement group comes online and joins the Lighthouse, it pulls the current model state from a healthy peer and rejoins training.

The critical property: if you have enough replica groups, a failure in one group does not require a training pause at all. The surviving groups continue advancing. The failed group rejoins when it recovers. PyTorch published a result showing a Llama 3 model trained across 300 L40S GPUs with synthetic failures injected every 15 seconds, continuing without restarts or rollbacks. (These are vendor-reported results from a controlled benchmark; results on real production hardware with real failure patterns will vary, but the architectural claim is sound.)

TorchFT works best when you can maintain at least a few replica groups in flight. This requires running more total GPUs than you strictly need for the model: if your model fits in 64 GPUs, you might run 4 replica groups of 64, giving you 256 GPUs total. That overhead is the cost of the checkpoint-less guarantee. For extremely long training runs on expensive frontier models, the trade-off often makes sense.

Within each replica group, standard FSDP2 or other intra-group parallelism is used. Across groups, TorchFT’s fault-tolerant DDP implementation handles gradient synchronization. The Lighthouse configuration drives quorum requirements: --min_replicas sets the minimum number of groups needed to proceed, --quorum_tick_ms controls how frequently membership is checked.

# Launch the Lighthouse server first
torchft_lighthouse --min_replicas 2 --quorum_tick_ms 100 --join_timeout_ms 10000 &

# Launch each replica group
torchrun \
  --nproc-per-node 8 \
  train_with_torchft.py \
    --replica_group_id 0 \
    --lighthouse_addr localhost:29510

TorchFT is still marked as an alpha prototype in the official repository as of this writing, and it carries the usual caveats about API stability. I would not run it on a production job without first validating it thoroughly on a representative workload.

Cloud-Specific Patterns: Where You Write Checkpoints Matters

The recovery architecture only works if you can write and read checkpoints fast. Distributed checkpointing produces I/O that scales with your GPU count: 256 GPUs each writing their own shard simultaneously. This can saturate standard NFS or POSIX file systems that were not designed for parallel write loads.

In practice this means:

For AWS, the combination of FSx for Lustre mounted to your EKS cluster is the right answer for most large-scale training. The parallel file systems article covers the throughput you can expect and how to tune the stripe configuration for parallel writes. For checkpoint reads on job restart, pre-loading shards into instance store NVMe before launching the job eliminates the checkpoint load time from your restart latency.

On Databricks AI Runtime, the UCVolumeWriter and UCVolumeReader backends stage through local NVMe and upload to Unity Catalog volumes asynchronously. This is a solid managed option if you are already in the Databricks ecosystem; it removes the infrastructure management overhead while giving you checkpointing nearly as fast as local NVMe.

On GCP, the custom AI silicon landscape adds another layer. Google Cloud’s elastic training demonstration on Cloud TPU (published July 2026) showed a training job recovering after a TPU worker was terminated without restarting from scratch. TPU pods have native checkpoint support through the JAX/MaxText stack; the patterns are analogous but the tooling differs from the PyTorch ecosystem.

The GPU infrastructure observability layer is also essential here: you want to detect XID errors before they cascade into a full process group failure. DCGM Exporter feeding Prometheus lets you alert on correctable memory error rates climbing above baseline, which is often a leading indicator of an impending hard fault on a specific GPU.

Choosing the Right Recovery Strategy

After running training infrastructure for various teams, my mental model for choosing a strategy breaks down roughly like this.

If your training runs are short, say under 24 hours, and your cluster is relatively stable, well-tuned async DCP checkpointing with TorchElastic as a backstop covers most cases. The operational complexity is low, the tooling is stable, and the resume semantics are straightforward. Point torchrun at your script, configure a Rendezvous backend, and make sure your data loader state is in the checkpoint.

If you are running very long pre-training jobs, weeks to months on large clusters, TorchFT’s checkpoint-less model starts making economic sense. The overhead of running spare replica groups is real, but the productivity cost of full restarts from a checkpoint hours old is also real. The AI FinOps calculus needs to account for both sides of that equation. TorchFT is the right answer for this regime once it reaches stable status.

For jobs on preemptible or spot GPUs, TorchElastic is the clearest win. Preemptions are predictable (you get a warning), the elastic restart absorbs them cleanly, and the cost savings from spot instances are substantial on GPU instance types. Pair TorchElastic with frequent async checkpointing and the preemption risk becomes manageable.

One thing I want to emphasize: the Clockwork vendor benchmark from March 2026 compared checkpoint-restart, TorchFT, and their own product on a 64 H200 GPU run and found checkpoint-restart outperformed TorchFT in their specific test. I am not dismissing that result, but vendor benchmarks require careful interpretation. The test conditions, failure injection patterns, checkpoint intervals, and model architecture all affect the outcome. Run your own benchmarks on your actual workload before committing to an architecture.

Comparison table showing checkpoint-restart vs TorchElastic vs TorchFT across the dimensions of goodput, operational complexity, spare capacity requirement, and suitability by job length and cluster type

The Data Pipeline Checkpoint Nobody Remembers

I have reviewed post-mortems for three separate training runs where teams resumed a job from a model checkpoint and discovered weeks later that their loss curves had an inexplicable anomaly. Two of the three turned out to be data pipeline corruption: the model resumed from step N, but the data loader resumed from step 0, silently doubling the exposure to the first portion of the training dataset.

Checkpointing the data pipeline is not glamorous and the APIs are less unified than model checkpointing. For torch.utils.data.DataLoader with a custom sampler, you need to save the sampler’s state so it can skip to the right position on resume. For streaming datasets and Mosaic’s StreamingDataset, the library provides native checkpoint support. For pipelines built on WebDataset or similar tar-based formats, you need to track which shards have been consumed and at what offset.

The RNG state is part of this too. Save torch.cuda.get_rng_state(), torch.get_rng_state(), and numpy.random.get_state() as part of your checkpoint. Without it, augmentations and dropout patterns diverge from what they would have been in the uninterrupted run.

These details feel minor until you are trying to explain to a research team why the model they trained for three weeks has behavior that does not match the original training plan.

Monitoring Goodput in Production

The metric I track on every training run is goodput: the ratio of time spent advancing the model (forward pass, backward pass, optimizer step) to total wall-clock time. Failures, checkpoint saves, and restarts all subtract from goodput.

With async checkpointing and TorchElastic, a well-run training job can maintain very high goodput even on commodity GPU hardware that sees occasional failures. The LLM fine-tuning infrastructure side of the problem, choosing the right parallelism strategy, gradient checkpointing trade-offs, and batch size, all feed into the goodput calculation. A training setup that is architecturally clean but loses three hours per failure because checkpoints are infrequent and slow can easily underperform a less elegant setup with disciplined async saves.

The GPU infrastructure observability stack from DCGM gives you the hardware signal. Prometheus and Grafana give you the goodput trend. Put them together and you can build dashboards that show, for each training run, exactly how much wall time was lost to each category of failure. That data drives your infrastructure investment decisions in a way that gut feeling about GPU reliability never will.

Pulling It Together

Fault-tolerant distributed training is not a single tool decision; it is a layered architecture. The layers are: fast parallel checkpointing via DCP as the foundation, async saves to make frequency cheap, data pipeline state as part of every checkpoint, TorchElastic for elastic restart on node failures, and TorchFT for checkpoint-less recovery on very long runs where spare replica groups make economic sense.

The GPU cluster networking infrastructure sets the ceiling on how fast collective communications can absorb a topology change when a node leaves and rejoins. The parallel file systems layer sets the ceiling on checkpoint throughput. The recovery strategy layer, TorchElastic versus TorchFT, sits on top of those and its value is proportional to how well the layers underneath it are tuned.

I have seen teams spend enormous effort optimizing their training code and then watch that effort evaporate in recovery time because they treated checkpointing as an afterthought. At the scale where training runs are measured in days and GPU-hours are priced accordingly, the reliability architecture is as important as the model architecture. Get both right.