Checkpointing and Mixed Precision
The One-Line Definition
Checkpointing saves a model's (and optimizer's) state to disk so training can be paused and resumed exactly where it left off, while mixed precision trains using a mix of 16-bit and 32-bit floating point numbers to cut memory usage and speed up training with minimal accuracy loss.
- Save and restore a fully resumable checkpoint (model, optimizer, scheduler, epoch, RNG state)
- Explain what autocast runs in reduced precision and what it keeps in float32
- Use GradScaler correctly for float16 and explain why bfloat16 usually doesn't need it
- Estimate the memory saved by mixed precision for a model
Checkpointing is like a video game save file for a model - if training crashes after 8 hours, or you just want to try the model as it stood at epoch 20, you can load that exact save point instead of starting over. Mixed precision is like using a smaller, faster notepad for most calculations while keeping a full-size notepad only for the parts where precision really matters - it lets you train bigger models faster on the same GPU without meaningfully hurting accuracy.
A checkpoint is a snapshot of model.state_dict() (all learnable parameters + buffers like BatchNorm running stats) plus optimizer.state_dict() (momentum/variance estimates for Adam, etc.) and any training metadata (epoch, best validation score, RNG state). Mixed precision (AMP - Automatic Mixed Precision) runs most ops in float16/bfloat16 while keeping a master copy of weights and precision-sensitive ops (reductions, loss computation) in float32, using a GradScaler to prevent gradient underflow in float16.
flowchart LR
subgraph CKPT["💾 Checkpointing"]
direction TB
S["model.state_dict()\n+ optimizer.state_dict()\n+ epoch/metadata"] --> Save["torch.save()"]
Save --> File["checkpoint.pt"]
File --> Load["torch.load()"]
Load --> Resume["load_state_dict()\nresume training"]
end
subgraph AMP["⚡ Mixed Precision"]
direction TB
FP32["🔵 Master weights\n(float32)"] --> Cast["autocast()\ncasts ops to fp16/bf16"]
Cast --> Fwd["Forward + Loss\n(fp16 compute)"]
Fwd --> Scale["GradScaler\nscales loss up"]
Scale --> Bwd["backward()\n(scaled grads)"]
Bwd --> Unscale["Unscale + step()\n(fp32 update)"]
end
style S fill:#d8dfe8,stroke:#b0bac8
style Save fill:#e8e0d4,stroke:#c8b89a
style File fill:#dde4dc,stroke:#b0c4b0
style Load fill:#e8e0d4,stroke:#c8b89a
style Resume fill:#ddd8e4,stroke:#b8b0c8
style FP32 fill:#d8dfe8,stroke:#b0bac8
style Cast fill:#e8e0d4,stroke:#c8b89a
style Fwd fill:#dde4dc,stroke:#b0c4b0
style Scale fill:#ddd8e4,stroke:#b8b0c8
style Bwd fill:#dde4dc,stroke:#b0c4b0
style Unscale fill:#d8dfe8,stroke:#b0bac8
Saving and Loading state_dict()
PyTorch recommends saving just the model's "settings" (its learned numbers), not the entire Python object - this is more portable and less likely to break when code changes slightly. To resume training properly, you also need to save the optimizer's internal state (things like momentum), otherwise resumed training can behave differently than an uninterrupted run.
# ---- Saving a checkpoint ----
checkpoint = {
"epoch": epoch,
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"best_val_loss": best_val_loss,
}
torch.save(checkpoint, f"checkpoints/model_epoch{epoch}.pt")
# ---- Loading and resuming ----
checkpoint = torch.load("checkpoints/model_epoch10.pt", map_location=device)
model.load_state_dict(checkpoint["model_state_dict"])
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
start_epoch = checkpoint["epoch"] + 1
best_val_loss = checkpoint["best_val_loss"]
model.to(device)
for epoch in range(start_epoch, num_epochs):
... # continue training exactly where it left off
torch.save(model.state_dict(), ...) is preferred over torch.save(model, ...) (pickling the whole object) because it's robust to refactors in the model's class definition, doesn't tie the checkpoint to exact file paths/class locations, and is the portable format expected by most downstream tooling. Always pass map_location=device when loading, so a checkpoint saved on GPU can still be loaded on a CPU-only machine.
Loading safely: since PyTorch 2.6, torch.load defaults to weights_only=True, which refuses to unpickle arbitrary Python objects (pickles can execute code). Keep checkpoints to tensors and plain Python types - as above - and load untrusted weights only from safetensors files (see PyTorch for LLMs).
Saving only the best model (common pattern to avoid disk bloat from saving every epoch):
if val_loss < best_val_loss:
best_val_loss = val_loss
torch.save(model.state_dict(), "checkpoints/best_model.pt")
Mixed Precision with autocast and GradScaler
GPUs can do 16-bit math roughly 2-3x faster than 32-bit math and it uses about half the memory - which means either training faster, or fitting a bigger model/batch size in the same GPU. The catch is that 16-bit numbers have a much smaller range, so very small gradient values can round down to zero ("underflow") and get lost. GradScaler works around this by temporarily multiplying the loss by a large number before the backward pass, then dividing back out before updating weights - keeping small gradients away from that underflow zone.
from torch.amp import autocast, GradScaler # torch.cuda.amp is deprecated
model = MyModel().to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
scaler = GradScaler("cuda") # only needed for float16, not bfloat16
for epoch in range(num_epochs):
model.train()
for batch_x, batch_y in train_loader:
batch_x, batch_y = batch_x.to(device), batch_y.to(device)
optimizer.zero_grad()
with autocast(device_type="cuda", dtype=torch.float16): # ops inside run in fp16 where safe
outputs = model(batch_x)
loss = loss_fn(outputs, batch_y)
scaler.scale(loss).backward() # scale loss up before backward - avoids fp16 underflow
scaler.step(optimizer) # unscales gradients, then calls optimizer.step()
scaler.update() # adjusts the scale factor for next iteration
Why this saves VRAM and speeds up training:
- Memory: activations and gradients stored in
float16(2 bytes) instead offloat32(4 bytes) roughly halves activation memory - the single biggest VRAM consumer during training for large batch sizes/sequence lengths - Speed: modern GPU tensor cores (NVIDIA Volta and newer) execute
float16/bfloat16matrix multiplications at significantly higher throughput thanfloat32 autocast()automatically chooses which ops are safe to run in reduced precision (matmuls, convolutions) versus which needfloat32for numerical stability (softmax, loss reductions, batch norm statistics) - you don't have to manually cast tensorsbfloat16(available on newer GPUs/TPUs) has the same exponent range asfloat32, so it often doesn't needGradScalerat all - it trades some mantissa precision for range stability, avoiding the underflow problemfloat16has
Study Notes
- Save
state_dict()(model + optimizer), not the whole model object, for portable and refactor-safe checkpoints - A resumable checkpoint needs: model weights, optimizer state, epoch number, and any tracked best-metric value
- Always pass
map_location=devicewhen loading a checkpoint to avoid device-mismatch errors autocast()picks safe ops to run in reduced precision automatically; you don't manually cast every tensorGradScalerprevents small gradients from underflowing to zero infloat16by scaling the loss up beforebackward()and unscaling beforeoptimizer.step()- Mixed precision roughly halves activation/gradient memory and speeds up compute on tensor-core GPUs, with minimal accuracy impact when done correctly
Check Yourself
- Why does bfloat16 training usually not need a GradScaler?
- Since PyTorch 2.6, what does torch.load do by default?
- Why save
model.state_dict()instead oftorch.save(model, path)? - What exactly does
GradScalerdo, step by step? - Why does
map_location=devicematter when loading a checkpoint? - Does mixed precision hurt model accuracy?
Exercises
Train for 3 epochs, save a checkpoint, train 2 more epochs and record the losses. Then restart from the checkpoint in a fresh process and train 2 epochs. Make the two loss sequences match exactly.
Hint
Save torch.get_rng_state(), torch.cuda.get_rng_state_all() and the DataLoader generator state, plus the scheduler state_dict
Solution
Matching requires model, optimizer and scheduler state, the epoch, all RNG states (Python, NumPy, torch CPU and CUDA) and the sampler's generator; any missing piece changes shuffling or dropout masks and the curves diverge. Deterministic kernels (torch.use_deterministic_algorithms) may also be needed on GPU.
Train the same model for 200 steps in float32, float16 with GradScaler, and bfloat16. Record peak memory (torch.cuda.max_memory_allocated), step time and final loss.
Solution
Both 16-bit runs cut activation memory roughly in half and are faster on tensor-core GPUs; final losses are close. float16 without GradScaler may show gradients underflowing (loss stalls); bfloat16 trains stably without scaling.
References
- PyTorch documentation, Automatic Mixed Precision package - torch.amp (2026)
- PyTorch documentation, torch.load (weights_only) (2026)
- PyTorch tutorials, Saving and Loading Models (2026)
- Micikevicius et al., Mixed Precision Training (2018)
- Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training (2019)
Last reviewed: 2026-09