Part of Language AI Handbook
Save and restore complete training state, choose checkpoint frequency, implement asynchronous I/O, and recover distributed LLM training runs from failures.
Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.
Article links
Make inline references clickable
Checkpointing and Recovery
Training a large language model takes weeks. A single job might run for twenty days, consuming thousands of GPU-hours and hundreds of thousands of dollars in compute. What happens when a GPU fails on day seventeen? Without a checkpoint, you restart from scratch. With a well-designed checkpointing strategy, you resume from a few hours ago at most, losing only a small fraction of your investment.
Checkpointing is the practice of periodically saving the complete state of a training run to durable storage, so that when something goes wrong (and in distributed training, something always goes wrong) you can restore that state and continue as if the interruption never happened. It sounds simple, but the implementation raises several questions: what exactly must you save, how often should you save it, how do you avoid stalling 512 GPUs while writing 200 GB to disk, and how do you ensure the recovered run produces the same results as the uninterrupted run would have?
The answer to that last question is worth lingering on. Exact reproducibility matters more than it might initially seem. During a long research run, you may need to debug a training divergence, verify that two experimental variants behaved identically up to a given step, or produce evidence for a paper that a particular result was obtained by a specific training procedure. A checkpointing strategy that recovers training in a way that looks approximately correct is not the same as one that recovers it exactly. The two goals require saving different state, and the cheaper approximation will eventually cost you.
This chapter covers checkpoint contents in depth, the tradeoffs behind checkpoint frequency decisions, asynchronous checkpointing to hide I/O latency, and the full fault recovery workflow that gets a distributed job back on its feet after a failure. It also covers the practical operational details that textbooks often skip: how to detect failures quickly, how to handle resharding when you recover onto different hardware, and how to defend against silent data corruption.
Building on the distributed training concepts examined in the data parallelism, tensor parallelism, and FSDP chapters, you now understand that training state is spread across many devices. Checkpointing must collect all of that distributed state into a coherent, recoverable snapshot. The next chapter covers monitoring and observability, which works hand-in-hand with checkpointing to detect failures quickly and trigger recovery.
What Goes Into a Checkpoint
A checkpoint must contain every piece of state that, if absent, would cause the resumed run to diverge from the original. This is a stricter requirement than it first appears.
The obvious items are model parameters and optimizer state. The less obvious items are the random number generator state, the data loader position, and the learning rate schedule step count. Miss any of these and your resumed run may converge differently, or produce non-reproducible results across restarts. Each component plays a specific role, and understanding that role helps you reason about what happens if it is omitted.
Model Parameters
Model parameters are the weights and biases of every layer. For a 70B parameter model in BF16, this is roughly 140 GB of raw tensor data. In a data-parallel setup, every rank holds a complete copy of the model, so in principle only one rank needs to write the checkpoint. In a model-parallel or FSDP setup, parameters are sharded across ranks, and the checkpoint may either consolidate shards into a single full-model file or store per-rank shard files.
Storing shards is faster: each rank writes only its slice independently, achieving bandwidth proportional to the number of ranks. If you have 64 ranks and 10 GB/s storage bandwidth per rank, the effective aggregate write bandwidth can approach 640 GB/s, cutting the checkpoint time from minutes to seconds. Consolidating to a single file is slower but simpler to load, especially when resuming with a different parallelism configuration. Most production systems use a middle ground: they write sharded checkpoints for fast recovery in place, and periodically consolidate a full-model checkpoint for portability and model release.
The consolidation step is usually done off the critical path, either on a scheduled basis or when a permanent checkpoint milestone is reached. A script reads all the per-rank shards and writes a single merged file. This merged file is what gets published, fine-tuned, or used for inference. The sharded checkpoints serve only the operational recovery use case.
Optimizer State
Optimizer state is often larger than the model parameters. For Adam, each parameter requires two additional tensors: the first-moment estimate and the second-moment estimate . Both tensors have the same shape as the parameter tensor they correspond to, and both are typically stored in FP32 (four bytes each) for numerical stability even when the model parameters are stored in BF16 (two bytes each). This means Adam optimizer state alone requires eight bytes per parameter, compared to two bytes for the BF16 parameter itself. The combined storage cost is ten bytes per parameter, roughly five times the parameter-only storage.
The optimizer state must be checkpointed alongside the parameters. Resuming with the parameters at step but with optimizer state reset to zero would cause the optimizer to behave as though it were just starting from scratch, leading to a transient spike in the effective learning rate and a loss curve that dips before recovering. The first-moment estimate captures the exponentially weighted average of past gradients, giving a smoothed gradient signal that prevents oscillation. The second-moment estimate captures the variance of past gradients, allowing adaptive step sizes for each parameter. Both take many steps to warm up from zero to useful values. Discarding them throws away potentially thousands of steps of accumulated curvature information, and the optimizer must rebuild that information from scratch before training can progress efficiently.
This is particularly damaging in the middle of a long run where the optimizer has accumulated information about the loss landscape. A fresh Adam optimizer at step 50,000 is a fundamentally different optimizer than one that has been running since step 1, even if both are given the same parameters and the same learning rate.
Gradient Scaler State
When training with mixed precision (covered in the Mixed Precision Training chapter), a gradient scaler maintains a dynamic loss scale factor that prevents underflow in FP16 gradients. The scaler multiplies the loss by before backpropagation, producing gradients that are large enough to be representable in FP16. After backpropagation, it divides the gradients by before the optimizer step, restoring them to the correct scale.
The value of is adapted during training. If gradients overflow (producing NaN or Inf), the scaler reduces by a factor and skips the optimizer step. If gradients remain finite for many consecutive steps, the scaler increases to use more of the FP16 dynamic range. The scaler also tracks how many consecutive steps have been free of overflow, using this to decide when to attempt a scale increase.
If you do not checkpoint the scaler state, the resumed run starts with the default initial scale rather than the converged scale, causing the scaler to go through its warmup phase again and potentially skipping optimizer updates due to overflow checks. The scaler state is tiny, typically just a single float and a few integers, but omitting it causes visible artifacts in the loss curve immediately after resuming. The loss will be flat for a few hundred steps while the scaler finds its footing, then resume its original trajectory.
Learning Rate Scheduler State
The learning rate scheduler tracks the current step count to implement schedules like cosine annealing or warmup followed by decay. The scheduler does not automatically infer the current step from the model weights; it maintains its own counter. The counter is what maps from step number to learning rate value.
For a cosine annealing schedule with a warmup period, the learning rate at step is:
where:
- : peak learning rate, reached at the end of the warmup period
- : minimum learning rate, approached at the end of the cosine decay
- : number of warmup steps
- : total training steps
- : current step count, stored in the scheduler state
If you restore the model and optimizer but not the scheduler, the step counter resets to zero. For a cosine schedule with a long decay, this means the resumed run may suddenly jump from a very low learning rate (late in the cosine decay) back to the initial warmup value. A sudden jump from to midway through training corrupts the training dynamics. The model may never fully recover to the trajectory it would have followed if uninterrupted, or it may recover only after several thousand steps of wasted compute.
Random Number Generator State
This is the subtlest item. Every GPU maintains its own RNG state, which is consumed by operations like dropout, weight initialization, and stochastic data augmentation. If you do not restore the RNG state at checkpoint time, the sequence of random numbers after resumption will differ from what it would have been in the uninterrupted run.
For most practical purposes this does not affect final accuracy, because different random seeds rarely change it. The loss curves of two training runs that start from the same checkpoint but use different random seeds will be nearly identical in expectation. But it does affect reproducibility. If you need to reproduce a training run exactly (for debugging a divergence, for example, where you want to determine whether the divergence is caused by a specific gradient update or by random variation), you must save and restore the RNG state on every device.
In PyTorch, each GPU has its own CUDA RNG state, separate from the CPU RNG state. A complete checkpoint saves both:
rng_states = {
"cpu": torch.get_rng_state(),
"cuda": torch.cuda.get_rng_state_all(), # list, one per GPU
}Restoring both of these at load time ensures that every dropout mask, every random data augmentation, and every stochastic operation unfolds in exactly the same order as it would have in the uninterrupted run.
Data Loader State
The data loader position determines which examples have been seen and which have not. Without it, resumption may either repeat examples already processed (wasting compute and potentially biasing the model toward examples that appear twice in the same epoch) or skip examples that have not yet been seen (creating an epoch that is effectively shorter than intended).
For large-scale training, datasets are typically pre-shuffled and streamed in order. Restoring the data loader state means recording the global step and reconstructing which shard files and which positions within those shards were being read. For a simple case where each step consumes exactly tokens from a sequential stream, you can reconstruct the position from the step number: after step with batch size , you have consumed tokens. You resume by skipping the first tokens of the stream.
This approach works well for deterministic, sequential data loading. It becomes more complex with streaming datasets or on-the-fly shuffling, which may require replaying the shuffle RNG state as well. Some frameworks sidestep this complexity by using a deterministic mapping from step number to data sample: a fixed global shuffle computed before training begins. Given the step number, you can always compute exactly which samples belong to that step without tracking any additional state beyond the step counter.
The cost of getting this wrong is not always obvious immediately. If your data loader resumes from the wrong position, you may train on a subtly different distribution of examples than intended. For large enough datasets, this is usually harmless. For smaller datasets where seeing the exact right examples matters, or for datasets with known challenging examples that you want to ensure are seen at the right training stage, it matters more.
Summary of Checkpoint Contents
A complete checkpoint contains:
- Model parameters: weights and biases for all layers
- Optimizer state: first and second moment estimates for each parameter (for Adam)
- Gradient scaler state: loss scale and overflow counters (for mixed precision training)
- Scheduler state: current step count and any internal counters
- RNG state: per-device random number generator state (for full reproducibility)
- Data loader state: global step index and dataset position
- Training metadata: step number, epoch, wall-clock time, configuration hash
The training metadata is not strictly necessary for exact resumption, but it is invaluable for bookkeeping and debugging. Knowing the wall-clock time helps you estimate how long the run took; the configuration hash lets you verify you are resuming with the same hyperparameters and not accidentally loading a checkpoint from a different experiment. The step number lets you verify that the recovery logic loaded the right checkpoint and that the scheduler counter is consistent with it.
Checkpoint Frequency
How often should you checkpoint? The answer depends on three competing factors: the cost of lost work if a failure occurs, the overhead of writing the checkpoint, and the storage cost of retaining multiple checkpoints. These three factors pull in opposite directions, and the optimal policy depends on your specific infrastructure.
The Cost-of-Failure Model
Suppose your training run takes total hours and you checkpoint every hours. If a failure occurs at a uniformly random time between two checkpoint events, the expected work lost is hours. Over the entire run, if there are failures, the expected total work lost is hours. Since longer runs experience more failures, the expected fraction of total work lost scales with and the failure rate.
For production-scale runs, the mean time between failures (MTBF) for individual GPUs or network switches is on the order of months, but a cluster of thousands of devices will experience failures far more frequently. The failures are not limited to GPU hardware faults: network switches fail, cooling systems fail, storage systems become unavailable, and job schedulers sometimes preempt running jobs. Each of these events terminates the training job and requires recovery.
With 1,000 GPUs each having an MTBF of 1,000 hours, the cluster as a whole has an effective MTBF of just 1 hour. With 10,000 GPUs, it drops to 6 minutes. Checkpointing every 10-15 minutes is common for very large runs precisely because of this arithmetic.
Formally, if the failure rate per device is (failures per hour) and you have devices, the cluster failure rate is approximately , assuming independent device failures. The expected time until the first failure is:
where:
- : number of devices in the cluster
- : per-device failure rate (failures per unit time, the reciprocal of MTBF)
- : cluster-wide failure rate, assuming independent failures
For a sensible checkpoint interval , you want:
so that the expected fraction of work lost per failure is small. As grows, shrinks, requiring shorter checkpoint intervals to maintain the same expected loss per failure. The implication is stark: a policy that worked well for a 100-GPU run will be inadequate for a 10,000-GPU run at the same failure rate per device, even if the model is the same size.
Failures are not always uniformly distributed in time. Hardware failures cluster after initial deployment (infant mortality), during thermal stress events, and during firmware updates. Practical checkpoint intervals should be conservative enough to handle bursts of failures during these higher-risk periods, not just average-rate failures.
The I/O Overhead Model
Checkpointing is not free. Writing a 300 GB checkpoint to NFS or cloud object storage at 10 GB/s takes 30 seconds. If you checkpoint every 10 minutes, you spend 30/600 = 5% of compute time blocked on I/O. At 1%, this overhead is acceptable and generally goes unnoticed. At 10%, it becomes a meaningful training slowdown that shortens the effective training budget.
Synchronous checkpointing, where all GPU work pauses while the checkpoint is written, is the simplest approach but has the highest overhead. Asynchronous checkpointing, described in the next section, overlaps I/O with computation to recover most of this overhead. But even with asynchronous I/O, the serialization phase (copying tensors from GPU memory to CPU memory) blocks training briefly. For a 300 GB checkpoint at 60 GB/s PCIe bandwidth, the serialization alone takes about 5 seconds.
The I/O model also interacts with distributed training. In a model-parallel or FSDP setup where parameters are sharded, each rank writes only its shard. If you have 64 ranks each writing 5 GB to separate storage volumes with 10 GB/s bandwidth each, the parallel write takes 0.5 seconds, far better than having a single rank write 320 GB sequentially. The storage infrastructure therefore determines the effective checkpoint bandwidth and directly influences how frequently you can checkpoint without unacceptable overhead.
Storage Cost
Retaining every checkpoint would be prohibitively expensive. A 70B model in BF16 with full Adam state requires roughly 420 GB per checkpoint. If you checkpoint every 10 minutes over a 20-day run, that is 2,880 checkpoints totaling over 1.2 petabytes of storage.
The practical solution is a rolling checkpoint policy: retain only the last checkpoints, deleting older ones as new ones arrive. to is typical for operational recovery. With a 15-minute checkpoint interval, this gives you 45-75 minutes of history, enough to recover from virtually any transient hardware failure.
Beyond the rolling window, many teams also retain a sparse set of permanent checkpoints at regular intervals (every 1,000 steps, for example) for research purposes. These allow comparison at different training stages, ablation studies, and emergency fallback if a silent data corruption is discovered days after it occurred. Permanent checkpoints are a small fraction of all checkpoints but provide a safety net that rolling-only policies cannot provide.
Silent data corruption is a real concern and motivates the permanent checkpoint policy. GPUs occasionally produce bit-flip errors in computation results without any hardware exception. The error may not cause an obvious crash but instead slowly degrades the model, producing weights that are slightly wrong and growing slightly more wrong with each contaminated gradient update. If you discover this 12 hours after it began, you need a checkpoint from before the corruption. A retention policy with at 15-minute intervals only gives you 75 minutes of history, which is not enough when the corruption has been propagating for half a day.
The typical approach is two-tier retention: rolling checkpoints for operational recovery (last , deleted aggressively), and permanent checkpoints at milestone steps for long-term fallback. Some organizations add a third tier: archival checkpoints kept indefinitely for models that reach certain loss thresholds or evaluation benchmarks. These represent training milestones worth preserving regardless of whether they are needed for recovery.
Asynchronous Checkpointing
The key insight behind asynchronous checkpointing is that writing a checkpoint to disk does not require the GPU to be involved. Once the checkpoint data has been copied from GPU memory to CPU memory, the GPU can continue training while the CPU handles the I/O.
To understand why this separation is possible, consider what a checkpoint really is at the moment it is created: a snapshot of tensor values that were valid at step . Once those values are safely in CPU memory, the GPU is free to advance to steps , , and so on. The CPU copy remains a valid representation of step- state regardless of what the GPU does afterward. The I/O operation is therefore completely independent of GPU computation, and the two can proceed in parallel.
This observation motivates a clean architectural separation: checkpointing logic consists of a fast "snapshot" phase that briefly blocks training (the GPU-to-CPU copy) and a slow "persist" phase that runs in the background (the CPU-to-disk write). Minimizing the snapshot phase minimizes the training stall, while the persist phase can take as long as necessary without affecting training progress.
The Synchronous Baseline
In synchronous checkpointing, the workflow at each checkpoint step is:
- Training pauses
- GPU tensors are transferred to CPU memory
- CPU writes tensors to disk
- Training resumes
Both steps 2 and 3 contribute to the stall. The GPU-to-CPU transfer alone can take several seconds for a large model (at roughly 60 GB/s PCIe 4.0 bandwidth, a 300 GB model takes about 5 seconds), and the disk write takes tens of seconds (at 10 GB/s NFS bandwidth, the same 300 GB takes 30 seconds). During all 35 of these seconds, hundreds of GPUs sit completely idle.
For a checkpoint interval of 10 minutes (600 seconds), this represents about 5.8% overhead from checkpointing alone. Over a 20-day run, you lose more than a full day of compute to synchronous checkpointing. The overhead scales linearly with model size and inversely with checkpoint interval, creating strong pressure to checkpoint less frequently, which in turn increases the expected work lost per failure.
Asynchronous checkpointing breaks this uncomfortable tradeoff. By running the disk write in the background, you pay only the GPU-to-CPU copy cost (the fast serialization phase) while the GPU continues training. The disk write proceeds concurrently and the GPU never waits for it.
Overlapping I/O with Computation
Asynchronous checkpointing breaks the write into two stages that overlap with ongoing training:
First, at the checkpoint step, the training loop snapshots the GPU tensors to CPU memory. This pinned-memory copy is fast, typically 2-5 seconds for a 300 GB model at roughly 60-100 GB/s PCIe bandwidth. Training then resumes immediately after this copy, without waiting for anything to be written to disk.
Second, a background thread or process continues transferring the CPU buffer to disk while training proceeds on the GPU. Since the CPU has a complete copy of the checkpoint state at the moment of the snapshot, it does not need the GPU again for I/O. The background writer can proceed at whatever rate the storage system supports.
The result is that the effective checkpoint overhead is reduced to the GPU-to-CPU copy time (a few seconds) rather than the full GPU-to-CPU plus CPU-to-disk time (tens of seconds). For the 300 GB model example above, asynchronous checkpointing reduces the training stall from 35 seconds to 5 seconds, an 85% reduction.
There is one important safety condition: the background write must complete before the next checkpoint is triggered. If the background writer is still writing the step-1000 checkpoint when it is time to checkpoint step-1100, you have two choices: wait for the previous write to finish before starting the new one, or start the new write anyway and accept that you briefly need two full checkpoint buffers in CPU memory. In practice, most systems wait, which means the effective checkpoint interval is bounded below by the write time. If your checkpoint write takes 30 seconds and you want to checkpoint every 10 minutes, asynchronous I/O works fine. If you want to checkpoint every 20 seconds, it does not help.
PyTorch's Asynchronous Checkpointing API
PyTorch provides torch.distributed.checkpoint with async support through AsyncCheckpointManager. The manager accepts a saving plan and executes it in a background thread, allowing training to proceed concurrently:
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict
# Snapshot state to CPU (fast, blocks training briefly)
state_dict = get_state_dict(model, optimizer)
# Launch async write (non-blocking, returns a future)
writer = dcp.FileSystemWriter("/checkpoints/step_1000")
future = dcp.async_save(state_dict, storage_writer=writer)
# Training continues here; the future resolves in backgroundThe future object can be waited on at a later step if needed, or simply allowed to run to completion in the background. If the process crashes before the write completes, the partial checkpoint is typically left on disk but is not used for recovery (the previous complete checkpoint is used instead, which is why the "latest checkpoint" pointer is updated only after all ranks have confirmed their writes are complete).
Memory Constraints
Asynchronous checkpointing requires holding two full copies of the model state simultaneously: one on the GPU (continuing training) and one in CPU pinned memory (being written to disk). For a 70B model with full optimizer state at ten bytes per parameter, this means roughly 700 GB of CPU RAM in addition to the GPU memory already in use.
Many clusters have sufficient CPU RAM. Modern servers with 8 GPUs often have 1-2 TB of system RAM. For a 70B model running on 64 H100 GPUs, each rank's shard is about 11 GB, and each host needs two copies of roughly 88 GB of tensors (for 8 GPUs per host), well within typical host memory. But for very large models or memory-constrained systems, the CPU RAM requirement can become a binding constraint.
In such cases, partial asynchronous checkpointing is used: save the optimizer state (the largest component) asynchronously while saving the model parameters synchronously, or vice versa. This reduces the peak CPU memory requirement proportionally, at the cost of a longer training stall than fully asynchronous checkpointing. The system designer must choose the right balance based on the available CPU RAM and the desired training stall duration.
Another memory-conscious approach is staged asynchronous checkpointing: copy one layer's parameters from GPU to CPU, then allow the background writer to drain that layer to disk, then proceed to the next layer. This keeps the CPU buffer small but serializes the I/O, losing some parallelism. It is most useful for single-GPU training where CPU RAM is the dominant constraint.
Coordinating Across Ranks
In distributed training, each rank must participate in checkpointing to ensure the complete model state is saved. The ranks must coordinate to ensure they all snapshot at the same step and that all shards are committed before a checkpoint is considered complete.
Coordination happens at two points. First, all ranks must reach the checkpoint step before any of them begins snapshotting. A global barrier at the checkpoint step ensures this. Second, after all ranks have launched their background writers, a second barrier (or a coordinator-rank acknowledgment protocol) ensures that the "latest checkpoint" pointer is updated only after all shards have been successfully written to durable storage.
If the second barrier is omitted and rank 0 updates the "latest checkpoint" pointer immediately after launching its background write, a process crash on rank 3 partway through its write would leave the checkpoint in an incomplete state with the pointer pointing to it. Recovery would attempt to load from this checkpoint, find that rank 3's shard is corrupt or missing, and fail. The coordinator-rank acknowledgment protocol prevents this by requiring all ranks to confirm success before the pointer is moved.
A typical coordination pattern uses a two-phase approach. In the first phase, all ranks write their shard data asynchronously. In the second phase, each rank reports completion to a coordinator (usually rank 0 or a dedicated checkpoint daemon), and the coordinator updates the metadata only after all ranks have reported. If any rank fails to report within a timeout, the checkpoint is considered incomplete and the previous pointer is retained.
Fault Recovery
A complete checkpoint is only half the system. The other half is the recovery logic that detects failures, loads the checkpoint, and resumes training with minimal disruption.
The checkpoint and recovery subsystems must be designed together. A checkpoint saved in a format that is slow to load creates a system where failures are infrequent but recovery is painful. A checkpoint format that loads quickly but is expensive to write creates a system where checkpointing overhead limits how often you can protect your work. Production training frameworks co-design these two halves, often using sharded file formats that write and load in parallel across many storage nodes, and maintaining a lightweight metadata store that makes "find the latest valid checkpoint" a constant-time operation rather than a directory scan.
Failure Detection
Distributed training frameworks use heartbeat mechanisms and timeouts to detect failures. When a process crashes, adjacent processes eventually time out waiting for communication and raise an exception. The NCCL collective communication library, which handles gradient synchronization in most GPU training setups, has its own timeout mechanism that triggers when a collective operation does not complete within a configurable window.
The time from failure to detection varies considerably. NCCL timeouts are typically set to 1,800 seconds (30 minutes) to tolerate temporary network congestion without false positives. A slow network switch, a brief NIC hiccup, or a transient memory allocation stall can all look like a hung collective operation from NCCL's perspective, and triggering a full job restart in response to a 5-second network blip would be extremely disruptive. The conservative default timeout prevents this.
But 30 minutes of hanging is an enormous waste. A job that checkpoints every 10 minutes and then hangs for 30 minutes after each failure is losing more time to failure detection than to failed work. Modern training setups are moving toward faster detection using in-band monitoring: a watchdog process that actively monitors process health, GPU utilization, and collective operation latency, and triggers recovery within seconds rather than waiting for NCCL to time out.
The watchdog approach requires careful calibration. A watchdog that is too sensitive will trigger false restarts during legitimate pauses (for example, during a checkpoint write or a data loading bottleneck). A watchdog that is too conservative is no improvement over NCCL's default timeout. The right setting depends on the specific training workload and the network characteristics of the cluster. Teams running production training often tune their watchdog parameters empirically based on the failure patterns they observe.
When the orchestrator (SLURM, Kubernetes, or a custom job manager) detects the failed job, it terminates all running ranks and begins the recovery sequence. Terminating all ranks is necessary even if only one has failed, because the surviving ranks are hung waiting for the failed rank's contributions to collective operations.
Recovery Workflow
The recovery workflow after a failure typically proceeds in a specific order. Understanding each step and its potential failure modes helps you build a reliable recovery system rather than one that fails silently.
The orchestrator detects the failure and terminates all ranks in the job. It then determines the most recent valid checkpoint by reading the checkpoint metadata (the "latest checkpoint" pointer file). It restarts the job with fresh processes on the available, non-failed hardware. Each new process initializes its model and optimizer to default values, then loads the checkpoint, restoring all state to the values at the checkpoint step. Training resumes from the checkpointed step.
The step "restart on available hardware" deserves elaboration. If a GPU has permanently failed (a common scenario with high-memory-bandwidth GPU memory errors), the replacement hardware may not be immediately available. The job may run with fewer GPUs than the original, which requires adjusting the parallelism configuration. In elastic training setups (discussed in the context of PyTorch Elastic), the restarted job may run with fewer GPUs than the original, adjusting the batch size or gradient accumulation steps accordingly to maintain the effective batch size.
Maintaining the effective batch size is important for training dynamics. If you trained with an effective batch size of 2,048 and you restart with 25% fewer GPUs, you have two choices: reduce the batch size to 1,536 (which may change training dynamics if you are near a threshold where batch size affects convergence) or increase gradient accumulation steps to compensate (which maintains the effective batch size at the cost of slower throughput). Most production systems prefer the gradient accumulation approach.
Loading State
Loading a checkpoint involves the reverse of saving: reading sharded files from storage, distributing the appropriate shards to each rank, and populating the model and optimizer state dictionaries.
The straightforward case is loading into the same parallelism configuration as the original: the same number of ranks, the same tensor parallelism degree, the same FSDP sharding strategy. In this case, each rank simply reads its own shard file and loads it directly into the appropriate tensors. This is fast and requires no redistribution.
The complex case arises when recovering onto different hardware with a different parallelism configuration. If the original run used 64 GPUs in a 4-way tensor parallel, 16-way data parallel configuration and you restart with 48 GPUs (after losing 16 to hardware failure), the shard layout has changed. PyTorch's distributed checkpoint format handles this through resharding: it reads all shards and redistributes the tensors according to the new sharding plan, without requiring the checkpoint to be re-saved in the new format.
Resharding is computationally inexpensive but communication-intensive: each rank must receive tensor slices from multiple source ranks to assemble its new shard. In practice, this takes a few minutes for large models and is acceptable as a one-time cost per recovery event.
The load_state_dict approach for a FSDP model using distributed checkpointing:
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
set_state_dict,
)
# After model and optimizer are created with new parallelism config:
state_dict = {"model": model.state_dict(), "optimizer": optimizer.state_dict()}
dcp.load(
state_dict=state_dict,
storage_reader=dcp.FileSystemReader("/checkpoints/step_1000"),
)
set_state_dict(
model,
optimizer,
model_state_dict=state_dict["model"],
optim_state_dict=state_dict["optimizer"],
options=StateDictOptions(full_state_dict=False), # keep sharded
)After loading, restore the scheduler, scaler, and data loader state from the metadata in the checkpoint before calling optimizer.step(). The order matters: the scheduler must be restored before the optimizer, because some scheduler implementations modify the optimizer's learning rate in their step() method, and you want to restore to the correct learning rate, not override it with an incorrect one.
Verifying Recovery
After loading a checkpoint, it is good practice to verify that the resumed run matches expectations before continuing for thousands of steps. A corrupted or incomplete checkpoint can cause training to silently diverge, wasting compute and potentially making the model worse. Catching this early is far cheaper than discovering it after 12 hours of additional training.
Common verification steps include:
- Confirm the global step counter matches the checkpoint metadata
- Run one forward pass and compare the loss to the logged loss at that step
- Verify the learning rate matches what the scheduler should produce at that step
- Check that gradient norms in the first resumed step are consistent with the run history
A mismatch in any of these suggests an incomplete or corrupt checkpoint, and it is better to roll back to an earlier checkpoint than to continue with a corrupted state.
Why run a forward pass and compare to logged loss? During a training run, the loss at each step is typically logged to a metrics system. After recovery, if you run the same batch of data (using the restored data loader position and RNG state) and compute the loss, you should get the same value as what was logged at that step. An exact match confirms that the model parameters and all randomness-affecting state have been correctly restored. Even a small discrepancy (beyond floating-point rounding) indicates something was missed in the checkpoint.
This kind of checkpoint integrity check is particularly important in long-running research jobs where silent errors can propagate undetected for hours before manifesting as anomalous evaluation metrics. The forward pass check is fast (a single batch, no backward pass needed), and the cost of skipping it is potentially many hours of corrupt training.
Some teams automate this verification step into the recovery logic itself: the recovery script runs the forward pass check, and if it fails, it automatically tries the previous checkpoint in the retention window. This removes the need for human intervention in most recovery scenarios and reduces mean recovery time from tens of minutes (waiting for a human to notice and act) to seconds.
Code Implementation
Let's build a practical checkpointing system that demonstrates the core concepts: saving a complete checkpoint, loading it for recovery, and measuring the overhead of synchronous versus simulated asynchronous checkpointing.
Setup and Imports
We will work with a medium-sized transformer model to make the storage and timing numbers realistic. The model is deliberately sized to produce measurable I/O times without requiring GPU hardware.
import torch
import torch.nn as nn
import torch.optim as optim
torch.manual_seed(42)
# Simulate a medium-sized model (like a small transformer block stack)
class SmallTransformer(nn.Module):
def __init__(self, vocab_size=32000, d_model=512, nhead=8, num_layers=6):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
encoder_layer = nn.TransformerEncoderLayer(
d_model=d_model,
nhead=nhead,
dim_feedforward=d_model * 4,
dropout=0.1,
batch_first=True,
)
self.transformer = nn.TransformerEncoder(
encoder_layer, num_layers=num_layers
)
self.head = nn.Linear(d_model, vocab_size)
def forward(self, x):
x = self.embedding(x)
x = self.transformer(x)
return self.head(x)
model = SmallTransformer()
optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10000)
scaler = torch.amp.GradScaler("cpu")Model parameters: 51,714,304 Parameter storage (BF16 equivalent): 206.9 MB
This gives a concrete sense of the scale. A model of this size has parameters, optimizer state, and gradients that together span several hundred megabytes even for a "small" 6-layer transformer. Production models are hundreds of times larger.
Building a Complete Checkpoint
A complete checkpoint function captures all the state described earlier: parameters, optimizer, scheduler, scaler, step number, and data loader position.
import os
from pathlib import Path
def save_checkpoint(
checkpoint_dir, model, optimizer, scheduler, scaler, step, data_seed
):
"""
Save a complete training checkpoint.
data_seed: integer representing the RNG seed + step, used to reconstruct
which data examples correspond to this training step.
"""
checkpoint_dir = Path(checkpoint_dir)
checkpoint_dir.mkdir(parents=True, exist_ok=True)
checkpoint = {
"step": step,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict(),
"scaler_state_dict": scaler.state_dict(),
"rng_state": torch.get_rng_state(),
"data_seed": data_seed,
"config": {
"d_model": 512,
"nhead": 8,
"num_layers": 6,
},
}
checkpoint_path = checkpoint_dir / f"checkpoint_step_{step:08d}.pt"
torch.save(checkpoint, checkpoint_path)
# Update the "latest" pointer atomically
latest_path = checkpoint_dir / "latest.txt"
tmp_path = checkpoint_dir / "latest.txt.tmp"
with open(tmp_path, "w") as f:
f.write(str(checkpoint_path))
os.replace(tmp_path, latest_path) # atomic on POSIX
return checkpoint_pathThe atomic update of latest.txt is important. If the process crashes while writing the checkpoint file itself, latest.txt still points to the previous valid checkpoint. Recovery logic always reads latest.txt first, so a partially written checkpoint file is never used. Without the atomic update (for example, if you updated latest.txt first and then wrote the checkpoint), a crash between those two operations would leave latest.txt pointing to a checkpoint that does not yet exist.
The os.replace() call is atomic on POSIX-compliant filesystems because it is implemented as a single rename system call. There is no window of time during which an observer could see latest.txt pointing to a nonexistent file. This guarantee holds for local filesystems and most distributed filesystems, but may not hold for all cloud object storage systems, which is why production systems sometimes use a dedicated coordination service (like etcd or Zookeeper) to manage the checkpoint pointer.
import os
import tempfile
import time
# Simulate a few training steps
train_data_seed = 12345
with tempfile.TemporaryDirectory() as tmpdir:
# Simulate training for a few steps
step = 1000
# Quick fake training step
dummy_input = torch.randint(0, 32000, (2, 64))
dummy_target = torch.randint(0, 32000, (2, 64))
optimizer.zero_grad()
output = model(dummy_input)
loss = nn.CrossEntropyLoss()(output.view(-1, 32000), dummy_target.view(-1))
loss.backward()
optimizer.step()
scheduler.step()
start = time.time()
ckpt_path = save_checkpoint(
tmpdir, model, optimizer, scheduler, scaler, step, train_data_seed
)
save_time = time.time() - start
ckpt_size_mb = os.path.getsize(ckpt_path) / 1e6
checkpoint_stats = {
"path": ckpt_path,
"size_mb": ckpt_size_mb,
"save_time_s": save_time,
"step": step,
}Checkpoint saved in 0.95s Checkpoint size: 620.7 MB Effective write speed: 651 MB/s Step: 1000
Loading a Checkpoint for Recovery
The recovery function loads the checkpoint and restores every piece of state. Notice how it mirrors the save function exactly, restoring each component in the reverse order it would be needed during training initialization.
def load_checkpoint(checkpoint_dir, model, optimizer, scheduler, scaler):
"""
Load the most recent valid checkpoint and restore all training state.
Returns the step to resume from, or 0 if no checkpoint exists.
"""
checkpoint_dir = Path(checkpoint_dir)
latest_path = checkpoint_dir / "latest.txt"
if not latest_path.exists():
print("No checkpoint found, starting from scratch.")
return 0, None
with open(latest_path, "r") as f:
ckpt_file = f.read().strip()
if not os.path.exists(ckpt_file):
print(f"Checkpoint file {ckpt_file} not found, starting from scratch.")
return 0, None
checkpoint = torch.load(ckpt_file, weights_only=True)
model.load_state_dict(checkpoint["model_state_dict"])
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
scheduler.load_state_dict(checkpoint["scheduler_state_dict"])
scaler.load_state_dict(checkpoint["scaler_state_dict"])
torch.set_rng_state(checkpoint["rng_state"])
step = checkpoint["step"]
data_seed = checkpoint["data_seed"]
return step, data_seed# Test the recovery
model_recovered = SmallTransformer()
optimizer_recovered = optim.AdamW(
model_recovered.parameters(), lr=1e-4, weight_decay=0.01
)
scheduler_recovered = optim.lr_scheduler.CosineAnnealingLR(
optimizer_recovered, T_max=10000
)
scaler_recovered = torch.amp.GradScaler("cpu")
with tempfile.TemporaryDirectory() as tmpdir:
# Save a checkpoint (re-run since tmpdir is new)
ckpt_path2 = save_checkpoint(
tmpdir, model, optimizer, scheduler, scaler, step, train_data_seed
)
start = time.time()
recovered_step, recovered_seed = load_checkpoint(
tmpdir,
model_recovered,
optimizer_recovered,
scheduler_recovered,
scaler_recovered,
)
load_time = time.time() - start
# Verify parameter equality
params_match = all(
torch.allclose(p1, p2)
for p1, p2 in zip(model.parameters(), model_recovered.parameters())
)
# Verify optimizer state equality
orig_lr = optimizer.param_groups[0]["lr"]
recov_lr = optimizer_recovered.param_groups[0]["lr"]
recovery_stats = {
"load_time_s": load_time,
"recovered_step": recovered_step,
"params_match": params_match,
"orig_lr": orig_lr,
"recov_lr": recov_lr,
"lr_match": abs(orig_lr - recov_lr) < 1e-10,
}Checkpoint loaded in 0.11s Resumed from step: 1000 Parameters match: True Learning rate match: True (lr=1.00e-04)
The params_match: True and lr_match: True outputs confirm that the recovery function restores the model to a state that is bit-for-bit identical to the original. This is the ground truth for a correct checkpoint implementation.
Measuring Checkpoint Overhead and Frequency Tradeoffs
Let's simulate the overhead tradeoff at different checkpoint frequencies and model sizes to understand when asynchronous checkpointing becomes critical.
# Simulate checkpoint timing for different model sizes
# Step time represents one training step (forward + backward + optimizer step)
step_time_ms = 800 # ~800ms per step for a large model on GPU
# Model sizes: number of parameters in billions
model_sizes_b = [1, 7, 13, 30, 70]
# Storage per parameter: BF16 (2 bytes) + Adam state (2 * 4 bytes) = 10 bytes
bytes_per_param = 10
# Assume 10 GB/s write bandwidth (NFS or cloud storage)
write_bandwidth_gbps = 10.0
checkpoint_data = []
for size_b in model_sizes_b:
size_bytes = size_b * 1e9 * bytes_per_param
size_gb = size_bytes / 1e9
write_time_s = size_gb / write_bandwidth_gbps
write_time_ms = write_time_s * 1000
for interval_steps in [50, 100, 250, 500, 1000]:
interval_time_ms = interval_steps * step_time_ms
overhead_pct = (write_time_ms / interval_time_ms) * 100
checkpoint_data.append(
{
"model_size_b": size_b,
"interval_steps": interval_steps,
"write_time_s": write_time_s,
"overhead_pct": overhead_pct,
}
) Model | Interval | Write Time | Overhead %
--------------------------------------------------
7B | 100 steps | 7.0s | 8.8%
7B | 500 steps | 7.0s | 1.8%
70B | 100 steps | 70.0s | 87.5%
70B | 500 steps | 70.0s | 17.5%The table makes the problem concrete. A 70B model writing to 10 GB/s storage requires 70 seconds per checkpoint. At a 100-step interval (about 80 seconds of training at 800ms per step), the synchronous overhead exceeds 87%. Asynchronous checkpointing is not a nice-to-have in this regime; it is a requirement for practical operation.
Visualizing Checkpoint Frequency vs. Overhead

Visualizing Expected Work Lost vs. Checkpoint Interval

The two charts together tell a complete story. The overhead chart shows that synchronous checkpointing at short intervals is infeasible for large models. The wasted-work chart shows that long intervals cause unacceptable work loss at scale. The only way to satisfy both constraints simultaneously is asynchronous checkpointing, which eliminates most of the overhead and allows short intervals without the associated I/O cost.
Rolling Checkpoint Policy
A rolling checkpoint manager retains only the last checkpoints, deleting older ones automatically. The implementation also maintains permanent checkpoints at specified intervals, giving long-term fallback protection.
class RollingCheckpointManager:
"""
Manages a rolling window of checkpoints, retaining only the last K.
Also maintains permanent checkpoints at specified intervals.
"""
def __init__(self, checkpoint_dir, keep_last=5, permanent_every=1000):
self.checkpoint_dir = Path(checkpoint_dir)
self.checkpoint_dir.mkdir(parents=True, exist_ok=True)
self.keep_last = keep_last
self.permanent_every = permanent_every
self.history = [] # list of (step, path) tuples
def save(self, model, optimizer, scheduler, scaler, step, data_seed):
is_permanent = step % self.permanent_every == 0
prefix = "permanent" if is_permanent else "rolling"
checkpoint = {
"step": step,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict(),
"scaler_state_dict": scaler.state_dict(),
"rng_state": torch.get_rng_state(),
"data_seed": data_seed,
"is_permanent": is_permanent,
}
path = self.checkpoint_dir / f"{prefix}_step_{step:08d}.pt"
torch.save(checkpoint, path)
if not is_permanent:
self.history.append((step, path))
# Delete old rolling checkpoints beyond keep_last
while len(self.history) > self.keep_last:
old_step, old_path = self.history.pop(0)
if old_path.exists():
old_path.unlink()
# Update latest pointer
latest_path = self.checkpoint_dir / "latest.txt"
tmp_path = self.checkpoint_dir / "latest.txt.tmp"
with open(tmp_path, "w") as f:
f.write(str(path))
os.replace(tmp_path, latest_path)
return path, is_permanent# Demonstrate rolling policy behavior
with tempfile.TemporaryDirectory() as tmpdir:
manager = RollingCheckpointManager(tmpdir, keep_last=3, permanent_every=500)
saved_records = []
for step_i in [100, 200, 300, 400, 500, 600, 700]:
path, is_perm = manager.save(
model, optimizer, scheduler, scaler, step_i, 42
)
exists = path.exists()
saved_records.append(
{
"step": step_i,
"permanent": is_perm,
"exists": exists,
}
)
# Check what's still on disk after all saves
existing_files = sorted(Path(tmpdir).glob("*.pt"))
for record in saved_records:
# Re-check existence (rolling may have deleted)
record["exists_after"] = any(
f"step_{record['step']:08d}" in str(p) for p in existing_files
)Step | Permanent | Retained ---------------------------------- 100 | no | deleted 200 | no | deleted 300 | no | deleted 400 | no | yes 500 | yes | yes 600 | no | yes 700 | no | yes Files still on disk: 4 permanent_step_00000500.pt rolling_step_00000400.pt rolling_step_00000600.pt rolling_step_00000700.pt
The output shows the rolling policy in action: steps 100 and 200 were deleted once step 700 was saved (exceeding the keep_last=3 window), but step 500 was retained permanently. A team investigating a regression discovered six days later can roll back to the step-500 permanent checkpoint even after hundreds of subsequent rolling checkpoints have been purged. The two-tier policy provides both operational recovery (rolling checkpoints for recent failures) and research-grade recovery (permanent checkpoints for long-term fallback).
Checkpoint Timing Breakdown
Let's measure the actual time breakdown between the CPU copy phase and the disk write phase to understand the async checkpointing opportunity precisely.
import io
def measure_checkpoint_phases(model, optimizer):
"""
Measure time for each phase of checkpointing to understand async opportunity.
"""
results = {}
# Phase 1: Serialize to in-memory buffer (simulates CPU copy from GPU)
buf = io.BytesIO()
t0 = time.perf_counter()
torch.save(
{
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
},
buf,
)
t1 = time.perf_counter()
results["serialize_ms"] = (t1 - t0) * 1000
results["size_mb"] = buf.tell() / 1e6
# Phase 2: Write from memory buffer to disk (simulates background I/O thread)
buf.seek(0)
with tempfile.NamedTemporaryFile(delete=True, suffix=".pt") as f:
t2 = time.perf_counter()
f.write(buf.read())
f.flush()
os.fsync(f.fileno())
t3 = time.perf_counter()
results["write_ms"] = (t3 - t2) * 1000
results["total_sync_ms"] = results["serialize_ms"] + results["write_ms"]
results["async_stall_ms"] = results[
"serialize_ms"
] # only stall during serialize
results["async_savings_pct"] = (
results["write_ms"] / results["total_sync_ms"]
) * 100
return resultstiming_results = measure_checkpoint_phases(model, optimizer)Checkpoint Timing Breakdown: Serialize (CPU copy equivalent): 200.7ms Write to disk: 208.6ms Total (synchronous): 409.3ms Async stall (serialize only): 200.7ms Time savings from async I/O: 51% Checkpoint size: 620.7MB
The timing breakdown confirms the core intuition: the disk write phase dominates total checkpoint time. Asynchronous checkpointing eliminates this dominant cost from the training stall, reducing the observable pause to just the serialization time. Even on a local SSD, disk writes take substantially longer than in-memory serialization for the same data, and the gap grows considerably on network-attached storage or cloud object stores.
Synchronous vs. Asynchronous Stall Comparison
To make the benefit of async checkpointing concrete across model scales, let's visualize how the synchronous stall (full serialize plus write) compares to the asynchronous stall (serialize only) for models ranging from 1B to 70B parameters.

The reduction for the 70B model is substantial. At that scale, the disk write alone consumes 70 seconds, while the serialization takes only about 11 seconds. Asynchronous checkpointing recovers those 70 seconds of idle GPU time entirely, at the cost of holding an extra copy of the model state in CPU memory.
Limitations and Practical Considerations
Checkpointing solves the fault tolerance problem elegantly at the conceptual level, but several real-world limitations shape how it is deployed in practice. Understanding these limitations helps you design a checkpointing strategy that works under production conditions rather than just in theory.
Storage Bandwidth and Cost
Large model checkpoints are expensive to store and transfer. A 70B model with full optimizer state requires roughly 420 GB per checkpoint. At cloud storage prices (approximately $0.02 per GB-month for standard storage), retaining 100 such checkpoints costs over $840 per month just for the storage. With fast object storage tiers, the per-GB price can be several times higher. Organizations must balance the safety of frequent checkpoints against their storage budgets, which often leads to aggressive rolling policies that minimize historical retention.
Storage bandwidth is equally constraining. Writing 420 GB at 10 GB/s takes 42 seconds. Even with asynchronous I/O, if your checkpoint interval is 10 minutes and your write takes 42 seconds, you need the async I/O to complete before the next checkpoint is triggered. For very large models or slow storage systems, you may need to extend the checkpoint interval, accept a larger I/O overhead, or invest in faster parallel storage. Some organizations solve this by provisioning high-bandwidth local NVMe storage on each training node, achieving 30-50 GB/s per node and dramatically shortening checkpoint write times.
The choice of checkpoint format also affects bandwidth requirements. PyTorch's torch.save() uses pickle serialization, which is flexible but not optimally efficient for large tensors. The safetensors format, developed specifically for large model weights, uses a simple binary layout that can be memory-mapped directly, enabling faster loads and more efficient storage. For very large models at scale, switching from torch.save() to safetensors can reduce both checkpoint write time and load time significantly.
Silent Data Corruption
Silent data corruption (SDC) is a failure mode where computation produces incorrect results without any obvious error signal. The model's loss may continue to decrease normally, but the corrupted values spread through the model parameters over subsequent steps. By the time the corruption is detected (often through downstream evaluation on held-out tasks), it may have propagated through hundreds of checkpoints and thousands of training steps.
SDC is more common on commodity GPU hardware than on ECC-protected memory, and its frequency increases with model size and training duration. Major LLM training efforts at companies like Google and Meta have reported SDC events that required careful forensic analysis to trace back to specific hardware faults. The events ranged from subtle bit flips that produced slightly wrong gradient values to more dramatic faults that corrupted entire parameter tensors.
The main defense against SDC is maintaining a long-term archive of permanent checkpoints with sufficient spacing to provide a clean restore point before the corruption began. Some teams also run periodic consistency checks: saving a "shadow" checkpoint and verifying that the model's forward pass produces the expected outputs on a fixed validation batch. If the forward pass loss drifts unexpectedly from the historical baseline, the consistency check triggers an alert before too much corrupted training has accumulated.
A complementary defense is gradient anomaly detection: tracking gradient norms per layer and alerting when any layer's norm exceeds a threshold multiple of its recent historical average. A single SDC event often produces an anomalously large gradient in the affected layer during the step when the corruption first affects the computation. Catching this in real time, before the next checkpoint is written, prevents the corrupt state from ever entering the checkpoint.
Checkpoint Consistency in Distributed Settings
In distributed training, all ranks must checkpoint at exactly the same step. If rank 0 checkpoints at step 1,000 but rank 1 checkpoints at step 1,001 (due to a scheduling discrepancy or a clock skew between nodes), the combined checkpoint is inconsistent: the model parameters and optimizer state from the two ranks are out of sync. Loading such a checkpoint and running a forward pass would produce different results on different ranks, causing collective operations (like AllReduce for gradient synchronization) to hang or produce garbage values.
Ensuring consistency requires a barrier at checkpoint time: all ranks must complete their writes before the "latest checkpoint" pointer is updated. PyTorch's distributed checkpoint library manages this through a coordinator rank that only marks the checkpoint complete after receiving acknowledgment from all participant ranks. The barrier is implemented using a distributed rendezvous operation that guarantees all ranks have committed their shard before any rank proceeds.
This consistency requirement interacts with the asynchronous checkpointing design. In the asynchronous case, the training barrier is placed after the fast serialization phase (GPU-to-CPU copy), not after the slow write phase. All ranks finish their serialization and collectively confirm they are ready before any rank launches its background writer. This ensures the checkpoint represents a consistent model state, even though the background writes proceed independently and at potentially different speeds.
Recovery Time Objective
In production, the recovery time from failure to resumed training matters as much as the checkpoint interval. If checkpointing is frequent but loading is slow (for example, loading a 420 GB checkpoint over a slow network), the total downtime per failure can be significant even when very little training was lost. For a checkpoint loading at 5 GB/s, a 420 GB checkpoint takes 84 seconds just for I/O. Add process startup time (tens of seconds for a large job), initialization, data pipeline warmup, and verification, and recovery can take 5-10 minutes even with a checkpoint from 15 minutes ago.
For a cluster with a 6-minute MTBF (10,000 GPUs at 1,000 hours MTBF per device), spending 5-10 minutes recovering from each failure means the cluster spends nearly as much time recovering as it does training. Reducing recovery time is as important as reducing checkpoint interval.
Pre-fetching checkpoints to local storage, using high-bandwidth interconnects for checkpoint loading, and structuring checkpoints as smaller sharded files that load in parallel are all techniques used to minimize recovery time. Some organizations maintain a local SSD cache on each training node: the most recent checkpoint is always available locally, allowing fast recovery from node-local failures and shortcutting the network bottleneck for other failure types.
The recovery time also depends on how quickly failures are detected. As discussed in the fault detection section, NCCL's default 30-minute timeout means that the total outage time for a single GPU failure is 30 minutes (detection) plus 5-10 minutes (recovery), totaling 35-40 minutes even with a 15-minute checkpoint interval. Faster detection dramatically reduces total downtime per failure, often more than any optimization to the checkpoint format or storage infrastructure.
Summary
Checkpointing is the fault tolerance mechanism that makes long distributed training runs practical. A complete checkpoint must contain the model parameters and optimizer state, plus the scheduler state, gradient scaler, RNG state, and data loader position. Missing any of these can cause the resumed run to diverge subtly from the original, corrupting research reproducibility and potentially wasting hundreds of GPU-hours before the problem is noticed.
Checkpoint frequency is governed by the tradeoff between expected work lost per failure and I/O overhead. Large clusters fail frequently, requiring frequent checkpoints, but large models are expensive to checkpoint. For a 10,000-GPU cluster with per-device MTBF of 1,000 hours, the expected time between cluster-level failures is just 6 minutes, demanding checkpoint intervals of that order. Synchronous checkpointing at such intervals would be completely infeasible for large models.
Asynchronous checkpointing resolves this tension by overlapping disk I/O with GPU computation. The training stall is reduced from the full checkpoint write time (tens of seconds) to just the GPU-to-CPU copy time (a few seconds), recovering the majority of the checkpointing overhead. The trade-off is an increased CPU memory requirement: you need space for two full copies of the model state simultaneously.
Rolling checkpoint policies balance recoverability against storage cost by retaining only the last checkpoints while preserving permanent checkpoints at sparser intervals for long-term fallback. The two-tier policy provides operational recovery for transient hardware failures and research-grade recovery for silent data corruptions that may not be detected for days.
Fault recovery involves detecting the failure, identifying the latest valid checkpoint, restarting the distributed job (possibly with resharding if hardware has changed), loading and restoring all state, and verifying that the resumed run matches expectations before continuing. With well-designed checkpointing and fast failure detection, a training job that loses a GPU on day seventeen can resume within minutes from a checkpoint taken at most 15 minutes earlier, losing only a tiny fraction of the total compute investment.
The next chapter on monitoring and observability covers how to detect training anomalies in real time, which works alongside checkpointing to catch problems quickly and determine exactly when and why a training run deviated from expected behavior.
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about checkpointing and fault recovery in distributed LLM training.
Checkpointing and Recovery Quiz
Reference
Citation details
Cite or share this article.
Continue with the full handbook
This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.
Explore Language AI HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

Comments
No comments yet. Be the first to share your thoughts!