Triaging the Loss Spike Shape Before Touching Code
When a long training run crashes or reaches a scheduled preemption, resuming from disk should be an invisible event on your telemetry dashboard. If the loss curve jumps upward on the very first resumed step instead, you are dealing with broken training state. Before modifying training scripts or retraining from scratch, note that saving a training checkpoint is fundamentally distinct from activation recomputation techniques like gradient checkpointing, which solely reduces peak activation VRAM during forward passes and holds no persistent state across run boundaries. In addition, distinguish resume anomalies from mid-run loss spikes caused by bad data batches or edge-of-stability dynamics: if the spike coincides precisely with step zero of your resumed process, the root cause lies in your serialization and restoration logic.
The geometric profile of the post-resume loss curve immediately isolates the category of missing state. Categorizing the trajectory over the initial 50 to 500 resumed steps cuts your debugging search space in half before you inspect a single line of code.
- Fast-decaying spike (recovers within 50 to 200 steps): The loss jumps sharply on step one but quickly trends back toward the pre-interruption baseline. This indicates that model weights and core optimizer states are intact, but transient runtime states - such as dataloader sample order, data shuffling, worker random number generator (RNG) seeds, or dynamic loss scalers - were reset and had to re-adapt.
- Persistent or diverging elevation (remains high or trends upward): The loss jumps and stays elevated for hundreds of steps, or gradients explode to infinity. This signals that the optimizer first- and second-moment buffers were omitted, sharded parameter mappings were corrupted, or the learning rate scheduler reset to step zero and applied a destructive update to existing weights.
- Immediate drop in loss (abnormally low loss): The loss plummets well below the expected trajectory. This indicates that the dataloader sampler reset to index zero without recording consumed samples, causing the model to overfit on data batches it has already memorized earlier in the epoch.
Missing Optimizer Moment State and Memory Accounting
The single most common bug in custom training loops is saving only the model state dictionary via torch.save(model.state_dict()) while omitting the optimizer state dictionary. For first-order optimizers like standard SGD without momentum, restoring weights alone is sufficient. However, modern large language model training relies almost universally on Adam or AdamW. These adaptive algorithms maintain running exponential moving averages of both past gradients (the first moment buffer, m_t) and squared gradients (the second raw moment buffer, v_t), alongside an internal step counter t used for analytical bias correction.
When you resume an AdamW optimizer from an unpopulated state, m_0 and v_0 are initialized to zero tensors. On the first resumed optimization step, the bias-corrected first moment estimator scales by 1/(1 - beta_1^t) with t=1, while the second moment denominator scales by sqrt(1/(1 - beta_2^t)). Because the historical variance buffer has been wiped out, the denominator fails to penalize high-variance coordinate directions. The resulting parameter updates are drastically misaligned in direction and orders of magnitude larger than what the local loss surface curvature can tolerate, destabilizing trained layers.
Saving these optimizer buffers introduces substantial disk and memory overhead that teams frequently overlook during infrastructure capacity planning. Depending on your training framework and numerical precision format, the total memory cost per parameter breaks down across distinct conventions:
| Framework / Setup | Accounting Convention | Bytes / Parameter | 70B Model Optimizer Overhead |
|---|---|---|---|
| DeepSpeed ZeRO-1 / ZeRO-2 | 2B fp16 weights + 2B fp16 grads + 4B fp32 master + 4B momentum + 4B variance ZeRO-3 vs FSDP | 16 bytes | 1,120 GB |
| Hugging Face Accelerate / Transformers | Standard fp32 master weights + fp32 gradients + 8B fp32 AdamW moments | 18 bytes | 1,260 GB |
| QLoRA / PEFT (bf16 base) | bf16 adapter weights + bf16 gradients + 8B fp32 AdamW moments (no fp32 master copy) | 12 bytes | 840 GB (on trainable baseline) |
Because the optimizer state for a 70B parameter model runs to hundreds of gigabytes under any of these conventions, saving incomplete checkpoints is often an accidental byproduct of trying to minimize I/O write latencies. However, omitting moment tensors completely invalidates the mathematical continuity of the optimization trajectory.
FSDP Shard Mismapping and World Size Changes
In distributed training using Fully Sharded Data Parallel (FSDP) or DeepSpeed ZeRO, optimizer states and model weights are partitioned across the data-parallel world size. A critical trap occurs when a run is saved on N GPUs and subsequently resumed on M GPUs, for example saving from a full node and restarting on a smaller or larger pool of devices. Sharded checkpoint files represent flat 1D tensor slices of partitioned parameters tied directly to rank IDs. If you attempt a naive rank-to-rank checkpoint load across mismatched world sizes, the optimizer buffers are silently assigned to incorrect parameter slices without throwing a runtime error, corrupting gradient updates immediately upon step execution.
Changing the world size also silently alters the global effective batch size unless compensation logic is explicitly added to the runtime configuration. Effective batch size is defined as the product of per-device micro-batch size, gradient accumulation steps, and the distributed data-parallel world size. If you reduce world size from 8 to 4 without doubling gradient accumulation steps, the effective batch size is halved. Under linear or square-root learning rate scaling rules, an optimizer operating with the original learning rate against a halved effective batch size is over-parameterized, causing severe loss oscillations.
To guarantee resharding resilience across dynamic GPU topologies, use PyTorch's native Distributed Checkpoint (DCP) framework via torch.distributed.checkpoint rather than legacy per-rank pickling. PyTorch's own recipe states that DCP saves and loads from multiple ranks in parallel and lets you re-shard across differing cluster topologies at load time, while its state_dict helpers manage fully qualified name mappings across models and optimizers:
- Wrap model and optimizer state dictionaries in stateful containers using torch.distributed.checkpoint.state_dict.get_state_dict.
- Save the distributed checkpoint directory concurrently across all ranks with dcp.save(state_dict, checkpoint_id=path).
- On resume, instantiate the model and optimizer under the new cluster world size and invoke dcp.load(state_dict, checkpoint_id=path) before training begins. DCP maps fully qualified parameter names (FQNs) to global tensors and redistributes tensor shards dynamically across the active ranks.
Learning Rate Scheduler and Step Counter Resets
Learning rate schedulers in PyTorch are step-counted state machines, not wall-clock timers. Schedulers like CosineAnnealingLR, LinearLR, or custom warmup-decay schedules track an internal last_epoch counter (which increments on every scheduler.step() call). If you instantiate a fresh scheduler during resume without loading its state dictionary via scheduler.load_state_dict(checkpoint['scheduler']), the step counter resets to zero.
- Warmup Schedule Reset: If your training run used a linear warmup phase and was interrupted deep into the cosine decay, resetting the scheduler re-engages the warmup curve. The learning rate collapses back to its tiny warmup floor, stalling parameter updates and degrading normalization statistics.
- Post-Warmup Decay Reset: If the scheduler resets to the peak learning rate on a model that has already reached late-stage convergence, the optimizer applies massive updates across low-curvature valleys. This blows the model out of its local minimum, resulting in a non-recoverable loss spike.
- Optimizer Calling Order Inversion: Calling optimizer.load_state_dict() after scheduler initialization without loading scheduler state can silently overwrite param_group learning rates back to initial config defaults.
To verify scheduler continuity during a resume event, implement a strict one-line parity check at the boundary of your training loop:
- Log the active learning rate from optimizer.param_groups[0]['lr'] immediately before serializing the checkpoint at step t.
- Log the active learning rate immediately after calling checkpoint restoration at resumed step t.
- Assert that abs(lr_resumed - lr_pre_crash) <= 1e-12. If the values differ by more than a single scheduler step increment, halt execution before computing forward passes.
Dataloader, Sampler, and RNG Reproducibility
Even when model weights, optimizer moments, and scheduler states are completely restored, training curves can diverge from numerical expectations if the data pipeline and pseudo-random number generators are not synchronized. In deep learning architectures utilizing stochastic regularization - such as Dropout, DropPath, FlashAttention dropout masks, and random sequence packing - output activations are determined by the internal states of active RNG engines.
For true deterministic continuity, your checkpointing routine must serialize and restore five distinct RNG state buffers across CPU and accelerator hardware:
- Python built-in random state: random.getstate() and random.setstate()
- NumPy random generator state: np.random.get_state() and np.random.set_state()
- PyTorch CPU RNG state: torch.get_rng_state() and torch.set_rng_state()
- PyTorch CUDA RNG state: torch.cuda.get_rng_state() and torch.cuda.set_rng_state() (or torch.cuda.get_rng_state_all() across multi-GPU ranks)
- Distributed Sampler Offset: The active dataset epoch and batch offset index inside your DistributedSampler or streaming iterable dataloader.
If the dataloader state is omitted, the sampler restarts from index zero. For token-streaming datasets, this causes the resumed run to consume identical batches that were already processed in earlier steps, leading to an artificial drop in loss followed by a spike when unfamiliar data is finally encountered. For map-style datasets with shuffle=True, omitting worker generator seeds causes the dataloader to generate a completely new shuffling permutation, disrupting curriculum schedules and gradient variance.
Hidden State in Mixed Precision Gradient Scalers
When training in standard IEEE 16-bit half-precision (fp16), the dynamic range of the numerical representation is constrained by its 5-bit exponent: normalized values run from roughly 6.10e-5 (2^-14) up to a maximum of 65,504. PyTorch's own AMP documentation is blunt about the consequence: gradient magnitudes too small for fp16 flush to zero (underflow), so the updates for those parameters are simply lost, which is why gradient scaling multiplies the loss by a scale factor before the backward pass. Loss scaling was introduced for exactly this reason in the foundational mixed precision work of Micikevicius et al. In practice, torch.cuda.amp.GradScaler dynamically scales loss values by a factor S before the backward pass and descales gradients prior to optimizer updates.
The dynamic GradScaler maintains internal persistent states: the current scale factor _scale, a consecutive non-overflow step counter _growth_tracker, and growth/backoff multipliers. Over thousands of iterations, the scaler dynamically adjusts S to ride the upper limit of the fp16 range. If you omit scaler.state_dict() from your checkpoint, the resumed run re-initializes S to its large default starting value. On the first resumed backward pass, this oversized scale factor pushes gradient values past the top of the fp16 range, generating inf or NaN values. The scaler responds by skipping optimizer.step(), halving S, and wasting training steps until the scale factor stabilizes back to the steady-state operating point.
Using bfloat16 (bf16) mixed precision eliminates the need for dynamic loss scaling entirely. The bf16 format allocates 8 bits to the exponent (matching IEEE fp32 dynamic range) and 7 bits to the mantissa, compared to fp16's 5 exponent bits and 10 mantissa bits. Because bf16 natively represents values down to 1.17e-38, gradient underflow is virtually eliminated without a scaler. However, this is an engineering trade-off: bf16 sacrifices numerical precision for range. For architectures highly sensitive to precision in master weight updates, fp16 with a properly serialized GradScaler remains necessary.
Proving Resume Equivalence and the Recovery Checklist
Do not wait for a multi-week cluster training failure to discover that your resume logic is flawed. Before launching production workloads, run a resume-equivalence verification test on a lightweight configuration (such as a 1B model or a short 500-step run). Train continuously from step 0 to step 200 while logging per-step training losses to a reference array L_continuous. Repeat the run from step 0, trigger a hard interruption and checkpoint serialization at step 100, restore all states, and train to step 200, recording L_resumed. Assert that for all steps t in [101, 200], abs(L_continuous[t] - L_resumed[t]) <= epsilon, where epsilon accounts for minor non-deterministic floating-point kernel noise.
Achieving bitwise-identical loss replay (epsilon = 0.0) requires exact hardware parity, matching CUDA and cuDNN versions, identical world sizes, and strict deterministic algorithm flags (torch.use_deterministic_algorithms(True)), which can reduce kernel throughput. For production engineering, agreement within numerical noise provides robust verification.
- Model Parameters: Serialized via a resharding-aware format such as torch.distributed.checkpoint (DCP) or safetensors with fully qualified parameter names.
- Optimizer State Dictionary: Complete first (m_t) and second (v_t) moment buffers, fp32 master weights (if applicable), and parameter group configurations.
- Optimizer Step Counter: The internal step counter t required for AdamW bias correction calculations.
- Learning Rate Scheduler State: scheduler.state_dict() capturing the active step counter and current learning rate.
- Dynamic Loss Scaler: scaler.state_dict() tracking the adapted scale factor and growth counters (mandatory for fp16).
- RNG State Ensembles: torch.get_rng_state(), torch.cuda.get_rng_state_all(), np.random.get_state(), and random.getstate().
- Dataloader and Sampler Position: Dataset epoch, global sample index offset, and worker worker_init_fn seeds.
- Infrastructure Metadata: Global world size, per-node GPU counts, and gradient accumulation settings.
If you encounter an unrecoverable checkpoint where optimizer states were lost in an ongoing production run, do not continue training at the pre-crash learning rate. Instead, execute a controlled mitigation strategy: implement a short linear re-warmup starting from a small fraction of the target learning rate to allow AdamW variance buffers to re-accumulate without destructive parameter jumps. If the loss fails to recover over that re-warmup window and the steps that follow it, roll back to an earlier valid checkpoint. When managing large-scale distributed training on European infrastructure, such as dedicated capacity on Lyceum Serverless Training or On-demand GPU VM clusters, validating your checkpoint serialization pipeline against this checklist ensures that hardware preemptions never compromise training convergence.