Contents
Map

02 · Prog Langs

Checkpointing & Mixed Precision

View as:

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.

Learning objectives 40 min
By the end of this page you will be able to:
  • 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 of float32 (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/bfloat16 matrix multiplications at significantly higher throughput than float32
  • autocast() automatically chooses which ops are safe to run in reduced precision (matmuls, convolutions) versus which need float32 for numerical stability (softmax, loss reductions, batch norm statistics) - you don't have to manually cast tensors
  • bfloat16 (available on newer GPUs/TPUs) has the same exponent range as float32, so it often doesn't need GradScaler at all - it trades some mantissa precision for range stability, avoiding the underflow problem float16 has

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=device when loading a checkpoint to avoid device-mismatch errors
  • autocast() picks safe ops to run in reduced precision automatically; you don't manually cast every tensor
  • GradScaler prevents small gradients from underflowing to zero in float16 by scaling the loss up before backward() and unscaling before optimizer.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

Check yourself
0 / 6 answered
  1. Why does bfloat16 training usually not need a GradScaler?
  2. Since PyTorch 2.6, what does torch.load do by default?
  3. Why save model.state_dict() instead of torch.save(model, path)?
  4. What exactly does GradScaler do, step by step?
  5. Why does map_location=device matter when loading a checkpoint?
  6. Does mixed precision hurt model accuracy?

Exercises

Exercise - Prove your checkpoint is resumable

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.

Exercise - Measure mixed precision

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

Last reviewed: 2026-09

⚡AI-assisted content - always verify, always explore multiple perspectives·