How to Save and Load PyTorch Transformer Model Checkpoints: Complete Implementation Guide
Save complete training snapshots—including model weights, optimizer state, loss history, and training metadata—using torch.save, then restore them with torch.load and rehydrate the model via load_state_dict() for inference or continued training.
The FareedKhan-dev/train-llm-from-scratch repository demonstrates production-grade checkpoint management for decoder-only transformer models. Saving and loading PyTorch transformer model checkpoints correctly requires preserving both the neural network parameters and the training context to ensure reproducible results. This guide walks through the exact implementation used in the source code, from serializing the full training state to reconstructing the model for text generation.
Saving Complete Training Checkpoints in PyTorch
During the training loop in scripts/train_transformer.py, the code constructs a comprehensive checkpoint dictionary that captures every component needed to resume training or analyze convergence. The implementation stores six critical keys:
model_state_dict: Contains the transformer’s learned parameters viamodel.state_dict()optimizer_state_dict: Preserves AdamW’s internal buffers and momentum viaoptimizer.state_dict()losses: Tracks the list of training losses collected at each steptrain_loss/dev_loss: Stores final evaluation metrics on train and validation splitssteps: Records the total number of training steps completed aslen(losses)
The code assembles these into a single dictionary and persists it using torch.save, automatically incrementing the filename suffix if the target path already exists:
checkpoint = {
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'losses': losses,
'train_loss': train_loss,
'dev_loss': dev_loss,
'steps': len(losses),
}
torch.save(checkpoint, modified_model_out_path)
Source: scripts/train_transformer.py, lines 44-55
Loading Checkpoints for Inference
The generation script in scripts/generate_text.py demonstrates the proper restoration procedure. Loading PyTorch transformer model checkpoints requires four distinct steps to ensure the model architecture matches the saved weights and runs on the correct device.
1. Deserialize the Checkpoint File
Use torch.load with map_location to restore the saved dictionary while mapping tensors to the target device (CPU or CUDA):
checkpoint = torch.load(model_path, map_location=torch.device(device))
Source: scripts/generate_text.py, line 21
2. Recreate the Model Architecture
Instantiate the Transformer class using the identical hyper-parameters defined in config/config.py. This ensures the architecture matches the saved state dictionary:
model = Transformer(
n_head=config["n_head"],
n_embed=config["n_embed"],
context_length=config["context_length"],
vocab_size=config["vocab_size"],
N_BLOCKS=config["n_blocks"],
)
Source: scripts/generate_text.py, lines 23-30
3. Restore the Saved Weights
Load the stored parameters into the freshly instantiated model using load_state_dict():
model.load_state_dict(checkpoint['model_state_dict'])
Source: scripts/generate_text.py, line 31
4. Configure for Inference
Set the model to evaluation mode to disable dropout and batch normalization updates, then transfer it to the target device:
model.eval().to(device)
Source: scripts/generate_text.py, line 32
Why This Checkpoint Pattern Is Recommended
Separating the model architecture from the parameter state provides several operational advantages. By recreating the Transformer class from configuration and only loading the state_dict, you guarantee that code changes to the model definition remain compatible with older checkpoints. Storing the optimizer_state_dict preserves momentum buffers and learning rate schedules, allowing training to resume exactly where it left off. Additionally, recording metadata like losses and steps enables reproducible convergence analysis without external logging systems.
Practical Implementation Examples
Complete Saving Workflow
# Inside scripts/train_transformer.py
import torch
# Assuming model, optimizer, losses, train_loss, dev_loss are defined
checkpoint = {
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"losses": losses,
"train_loss": train_loss,
"dev_loss": dev_loss,
"steps": len(losses),
}
torch.save(checkpoint, "checkpoints/transformer_latest.pt")
Complete Loading and Generation Workflow
# Inside scripts/generate_text.py
import torch
import tiktoken
from src.models.transformer import Transformer
from config.config import config
def generate_text(model_path, input_text, max_new_tokens=100, device="cpu"):
# Load checkpoint with device mapping
checkpoint = torch.load(model_path, map_location=torch.device(device))
# Re-instantiate model architecture
model = Transformer(
n_head=config["n_head"],
n_embed=config["n_embed"],
context_length=config["context_length"],
vocab_size=config["vocab_size"],
N_BLOCKS=config["n_blocks"],
)
# Restore weights and configure for inference
model.load_state_dict(checkpoint["model_state_dict"])
model.eval().to(device)
# Tokenize input and generate
enc = tiktoken.get_encoding("r50k_base")
context = torch.tensor(
enc.encode_ordinary(input_text),
dtype=torch.long,
device=device
).unsqueeze(0)
with torch.no_grad():
generated = model.generate(context, max_new_tokens=max_new_tokens)[0]
return enc.decode(generated.tolist())
Summary
- Comprehensive state saving: Store
model_state_dict,optimizer_state_dict, loss curves, and step counts to enable full training resumption - Architecture recreation: Always instantiate the
Transformerclass fromconfig/config.pybefore loading weights to maintain compatibility - Device-agnostic loading: Use
map_locationintorch.loadto seamlessly transfer checkpoints between CPU and GPU environments - Evaluation mode: Call
model.eval()after loading to disable stochastic layers like dropout before inference - Source files: Implement saving in
scripts/train_transformer.pyand loading inscripts/generate_text.pywith architecture defined insrc/models/transformer.py
Frequently Asked Questions
What should be included in a PyTorch transformer checkpoint?
A production-ready checkpoint should contain the model_state_dict for weights, optimizer_state_dict for training resumption, loss history for monitoring convergence, and metadata like step count and evaluation metrics. The FareedKhan-dev/train-llm-from-scratch repository stores these in a dictionary keyed as model_state_dict, optimizer_state_dict, losses, train_loss, dev_loss, and steps.
Why recreate the model architecture instead of pickling the entire model?
Recreating the Transformer class from configuration and loading only the state_dict ensures forward compatibility with code changes. Pickling the entire model object risks breakage if the class definition changes, while the state dict approach decouples the architecture definition in src/models/transformer.py from the saved parameters.
How do I load a checkpoint on CPU when it was saved on GPU?
Pass map_location=torch.device('cpu') to torch.load when deserializing the checkpoint. This maps all tensors to the CPU device during loading, preventing CUDA availability errors and allowing inference on hardware without GPUs.
Can I resume training from a saved checkpoint?
Yes. To resume training, load the checkpoint using torch.load, recreate the model and optimizer, then call optimizer.load_state_dict(checkpoint['optimizer_state_dict']) and restore the step counter from checkpoint['steps']. This restores the optimizer's momentum buffers and learning rate state, allowing training to continue exactly from the saved iteration.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →