Skip to content

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

  1. 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)

  2. Monolithic torch.save on 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)

  3. Distributed Checkpoint (DCP) inverts the design: every rank writes its own shard in parallel, alongside a small .metadata file describing how shards compose into full tensors. Save time drops roughly as 1/N with rank count, and because .metadata records 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)

  4. 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)

  5. 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_save splits 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 for torch.save):

  6. DDP LLM, 2.8B params on 32×H100: 1.8× (36s vs 66s)
  7. 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)

  1. 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)

  2. 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 .metadata file 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)

  3. 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)

  4. 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. UCVolumeDataset streams 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 Runtime DataLoader is a drop-in PyTorch DataLoader subclass whose defaults are tuned for concurrent fetch-and-cache during compute. Example (image classification, per GPU, steady state): stock PyTorch reading directly from UC vs UCVolumeDataset + Databricks DataLoader:

  5. Epoch 1 throughput: 57.2 → 417 images/sec
  6. Epoch 2 throughput: 371.6 → 6,590 images/sec
  7. 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)

  1. 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

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_save numbers exclude network-storage time for the torch.save baseline (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 UCVolumeDataset partitioning, cache-eviction, and prefetch-depth parameters are not disclosed here.

Source

Last updated · 766 distilled / 2,225 read