Fast, fault-tolerant PyTorch training on AI Runtime¶
Summary¶
A Databricks AI Runtime engineering post arguing that large-scale training
efficiency reduces to a single metric — goodput, the
fraction of GPU time spent on productive computation rather than waiting or
recovering from failures — and that two routinely-neglected subsystems make or
break it: the checkpointing mechanism and the data pipeline. Because GPU
failures are the expected case at scale, fast automatic recovery is the only
way to keep goodput high and GPU spend bounded. The post walks four compounding
decisions: use PyTorch Distributed Checkpoint (DCP)
instead of monolithic torch.save (even for data-parallel jobs); make saves
asynchronous so frequency is nearly
free; recover automatically to
the most recent complete checkpoint; and overlap
dataloading with compute while caching remote data to local NVMe. A fifth,
subtler point: checkpoint the data-pipeline position and RNG state, or restarts
silently corrupt the training distribution.
On AI Runtime these are implemented by UCVolumeWriter/UCVolumeReader and
UCVolumeDataset + a tuned DataLoader over
Unity Catalog Volumes.
Key takeaways¶
-
Goodput is the master metric; failures are the expected case. At scale, the probability a job survives its full duration without interruption falls rapidly. Using a ~1% annualized per-GPU failure rate: a 256-GPU job over 30 days has ~19% chance of a failure; at 1,024 GPUs it climbs to ~57% — and these are only infrastructure-level issues. Grounding it in reality, the 608×H100 Delta supercomputer saw failures every 1.9 hours, implying a 32-GPU job would average ~36 hours to failure. Your job will fail; engineering choices set how much time you lose. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
-
Monolithic
torch.saveon rank 0 is a serial bottleneck. For distributed training it gathers all state to rank 0 and writes a single file; a single process writes synchronously and can block on network transfer to remote object stores (like Unity Catalog volumes). That blocking leaves GPUs idle, eroding goodput. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime) -
Distributed Checkpoint (DCP) inverts the design: every rank writes its own shard in parallel, alongside a small
.metadatafile describing how shards compose into full tensors. Save time drops roughly as 1/N with rank count, and because.metadatarecords the global layout, the same checkpoint can reload onto a different number of GPUs — DCP re-plans which bytes each new rank needs, so recovering onto a reduced-capacity cluster after losing nodes "just works." (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime) -
DCP is worth it even for plain data-parallel (DDP) jobs. A common assumption is that DCP only helps sharded models, since DDP replicas hold identical weights. Not so: DCP shards the model state and writes it in parallel even for DDP. It's also the same API you'll need the day you move to FSDP or tensor parallelism, so adopting it early means never rewriting resilience code "at the worst possible time." (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
-
Asynchronous saves make frequency nearly free. A synchronous save blocks training until bytes are durable — tens of seconds of idle accelerator time for a large checkpoint to a remote volume.
async_savesplits it: a fast copy to a staging buffer, then a background upload that overlaps continued training. The training loop pays only for the staging copy, not the upload. Measured savings (excluding network-storage time fortorch.save): - DDP LLM, 2.8B params on 32×H100: 1.8× (36s vs 66s)
- FSDP LLM, 20B params on 32×H100: 58× (522s vs 9s)
On AI Runtime, UCVolumeWriter/UCVolumeReader implement DCP against UC
volumes, staging I/O through local NVMe and marking a checkpoint complete
only once its data has fully landed. (Source:
sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
-
Checkpoint frequency sets recovery cost, and cheap async saves let you make the interval small. Expected wasted work per failure is ~half the checkpoint interval; cutting the interval 10× cuts expected recovery time 10×. Using Llama 3's cited ~8.6 interruptions/day: checkpointing every 2 hours wastes ~8.6 h/day → 64% goodput; checkpointing every 30 minutes wastes ~2.15 h/day → 91% goodput. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
-
Recovery must be automatic and select only complete checkpoints. On restart the job should find the most recent checkpoint that finished writing, skipping any left half-written by the crash, with no human in the loop. DCP makes this reliable: the
.metadatafile is written only after all shards land, so its presence is a trustworthy "this save is complete" marker to select on. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime) -
Dataloading decides whether GPUs ever idle. A job runs at the speed of its slowest input; when accelerators wait on the next batch, goodput drops (see GPU stall from storage). The fix is to overlap next-step data preparation with current-step compute. Databricks reports customers who shift to overlapped dataloading see a 20–50% decrease in wall-clock time. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
-
Cache remote data to local NVMe on first access. On a governed platform training data lives in remote object storage; UC volumes are surfaced as network mounts, and reading directly on every access binds step time to network latency and re-downloads the same files every epoch.
UCVolumeDatasetstreams files from a UC volume, caches each to local NVMe on first access, and partitions files across ranks and workers so every accelerator gets a disjoint, non-overlapping slice. The AI RuntimeDataLoaderis a drop-in PyTorchDataLoadersubclass whose defaults are tuned for concurrent fetch-and-cache during compute. Example (image classification, per GPU, steady state): stock PyTorch reading directly from UC vsUCVolumeDataset+ DatabricksDataLoader: - Epoch 1 throughput: 57.2 → 417 images/sec
- Epoch 2 throughput: 371.6 → 6,590 images/sec
- GPU utilization: 12.6% → 53.3%
The DataLoader also logs fetch_seconds (time to produce a batch = GPU idle
time) to MLflow so a blocking data pipeline is visible at
a glance. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
- Forgetting the data pipeline silently corrupts your model — the worst failure because there's no error, no crash, no failed job. If you checkpoint model + optimizer + step but not the dataloader's position, a mid-epoch restart replays already-seen examples and skips unreached ones, biasing the data distribution across the many restarts scale makes routine. The fix treats data position as part of the checkpoint (sample/shard offset, or epoch-boundary checkpointing) and rests on one prerequisite: determinism — shuffling/augmentation RNG seeds and states must be checkpointed too, or a saved position points at the wrong samples after restart. Seed, reproducible order, and resumable pipeline are three expressions of one idea. (Source: sources/2026-08-28-databricks-fast-fault-tolerant-pytorch-training-on-ai-runtime)
Systems / concepts extracted¶
- Systems: PyTorch (DCP,
async_save,DataLoader), PyTorch Distributed Checkpoint (DCP), AI RuntimeUCVolumeWriter/UCVolumeReader/UCVolumeDatasetover Unity Catalog Volumes, MLflow (dataloader metric logging), Databricks AI Runtime. - Concepts: goodput, distributed checkpoint, checkpoint frequency, data-pipeline checkpointing, GPU training failure modes, GPU stall from storage, annualized failure rate, FSDP.
- Patterns: async checkpoint staging, overlapped dataloading, local NVMe cache for remote training data, resume to latest complete checkpoint.
Operational numbers¶
| Fact | Value |
|---|---|
| Assumed per-GPU annualized failure rate | ~1% |
| 256-GPU / 30-day job failure probability | ~19% |
| 1,024-GPU / 30-day job failure probability | ~57% |
| Delta supercomputer (608×H100) MTBF | failure every ~1.9 h |
| Implied 32-GPU MTBF | ~36 h |
async_save vs torch.save, DDP 2.8B on 32×H100 |
1.8× (36s vs 66s) |
async_save vs torch.save, FSDP 20B on 32×H100 |
58× (522s vs 9s) |
| Llama 3 interruption rate | ~8.6 / day |
| Goodput @ 2h checkpoint interval | 64% (waste ~8.6 h/day) |
| Goodput @ 30min checkpoint interval | 91% (waste ~2.15 h/day) |
| Overlapped dataloading wall-clock reduction | 20–50% |
| Image workload epoch-1 throughput (stock → UCVolumeDataset) | 57.2 → 417 img/s |
| Image workload epoch-2 throughput | 371.6 → 6,590 img/s |
| Image workload GPU utilization | 12.6% → 53.3% |
Caveats¶
- Benchmarks are Databricks' own on AI Runtime; the
async_savenumbers exclude network-storage time for thetorch.savebaseline (so the true gap favouring async is understated). Image-throughput numbers are a single workload (JPEG decode + augment + vision model) on unspecified GPU/model/batch beyond "same GPU, model, and batch size." - The 1% annualized failure rate and the failure-probability curve come from a companion Databricks post's back-of-the-envelope model, not a measured fleet distribution.
- The post is a mechanisms-and-tradeoffs walkthrough with a companion
[performance-and-resiliency guide] for code; exact
UCVolumeDatasetpartitioning, cache-eviction, and prefetch-depth parameters are not disclosed here.
Source¶
- Original: https://www.databricks.com/blog/fast-fault-tolerant-pytorch-training-ai-runtime
- Raw markdown:
raw/databricks/2026-08-28-fast-fault-tolerant-pytorch-training-on-ai-runtime-f2d9d494.md
Related¶
- concepts/goodput — the master metric this whole post optimizes.
- distributed-checkpoint — the checkpoint format that unblocks frequency.
- checkpoint-frequency — the goodput lever cheap saves unlock.
- concepts/durable-execution — the silent-corruption failure mode.
- gpu-training-failure-modes — why recovery is the expected case.
- gpu-stall-from-storage — the dataloading failure mode.
- systems/pytorch-distributed-checkpoint — the DCP API and
.metadatamarker. - systems/unity-catalog-volumes — the governed remote store DCP/data load against.
- companies/databricks — the source company.