Saving and Loading PyTorch Model Checkpoints: A Practical Guide
Learn how to save and resume PyTorch training with torch.save and torch.load, handle device mapping, and verify checkpoint integrity.
29 Jul 2026, 22:30 UTC

Desired outcome
You want to persist a training session so you can stop, later resume, or evaluate the model on the same hardware or a different device. The guide shows how to save a checkpoint that includes model weights, optimizer state, epoch index and loss, and how to reload it safely with correct device mapping.
Prerequisites
- PyTorch installed (tested with 2.x series, but the steps work with any recent 1.x release).
- A training script that defines a
model(nn.Module), anoptimizer(e.g., SGD or Adam), and variablesepochandloss. - Write permission to the directory where the checkpoint file will be stored.
- If you plan to load on a GPU, ensure CUDA‑visible devices are available.
Procedure
1. Save a checkpoint
Place the saving block at the end of each epoch or whenever you need a snapshot. The dictionary contains everything required to resume training.
# Save checkpoint (run in your training loop, same process that owns model & optimizer) checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'loss': loss } torch.save(checkpoint, 'checkpoint.pt')Where to run: inside the Python process that executes the training script. Required permissions: ability to create/overwrite
checkpoint.ptin the current working directory or a specified path. Risks: the file will be overwritten each call; consider adding epoch number to the filename if you need a history.2. Load a checkpoint
When you restart training or run inference, load the file and transfer the tensors to the target device. Use
map_locationto avoid device‑mismatch errors.# Load checkpoint (run in a fresh Python process or after re‑initializing model/optimizer) checkpoint = torch.load('checkpoint.pt', map_location='cpu') # change to 'cuda:0' if loading directly to GPU model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) start_epoch = checkpoint['epoch'] + 1 last_loss = checkpoint['loss']Where to run: any Python environment with PyTorch available. Required permissions: read access to
checkpoint.pt. If you intend to train further on GPU, after loading callmodel.to('cuda')andoptimizertensors will follow automatically because they reference model parameters.Expected checks
- Verify the file exists:
import os; assert os.path.isfile('checkpoint.pt') - Confirm that tensors are on the expected device after loading, e.g.,
assert next(model.parameters()).device.type == 'cpu'(or 'cuda' if you moved them). - Spot‑check a parameter to ensure values match the original saved state:
original = torch.load('checkpoint.pt', map_location='cpu')['model_state_dict']['layer1.weight']; loaded = model.state_dict()['layer1.weight']; assert torch.allclose(original, loaded) - Optionally compute a simple hash of the state dict before saving and after loading to detect silent corruption.
Recovery options
If loading raises RuntimeError: Expected all tensors to be on the same device, you likely omitted or mis‑specified map_location. Reload with the correct device string.
If you see RuntimeError: Invalid pickle protocol or a version‑related error, the checkpoint was created with a different PyTorch major release. In that case, try re‑saving the checkpoint with the current PyTorch version, or load with pickle_module=<older version> if you must retain the old file.
Should the file become corrupted (zero bytes or I/O error), restore from a backup copy or restart training from the last known good checkpoint.
Limitations and practical verification
Checkpoint size grows with the number of parameters and optimizer state. To reduce disk usage, you can enable the newer zipfile serialization: torch.save(checkpoint, 'checkpoint.pt', _use_new_zipfile_serialization=True). This does not affect loading semantics.
Always verify that the reloaded model produces the same forward pass output for a given input as before saving (in evaluation mode). A quick test:
model.eval() with torch.no_grad(): out_before = model(sample_input) # re‑load as shown above model.eval() with torch.no_grad(): out_after = model(sample_input) assert torch.allclose(out_before, out_after, atol=1e-6)This confirms that the serialization‑deserialization round‑trip preserved the model’s behavior.
0 replies
A thoughtful contribution can make all the difference. Be the first to share one.