# How to Save and Load PyTorch Transformer Model Checkpoints: Complete Implementation Guide

> Learn to save and load PyTorch transformer model checkpoints including weights and optimizer state. Implement complete training snapshots using torch.save and torch.load for seamless restoration.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: how-to-guide
- Published: 2026-06-01

---

**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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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 via `model.state_dict()`
- `optimizer_state_dict`: Preserves AdamW’s internal buffers and momentum via `optimizer.state_dict()`
- `losses`: Tracks the list of training losses collected at each step
- `train_loss` / `dev_loss`: Stores final evaluation metrics on train and validation splits
- `steps`: Records the total number of training steps completed as `len(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:

```python
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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py), lines 44-55

## Loading Checkpoints for Inference

The generation script in [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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):

```python
checkpoint = torch.load(model_path, map_location=torch.device(device))

```

*Source:* [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py). This ensures the architecture matches the saved state dictionary:

```python
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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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()`:

```python
model.load_state_dict(checkpoint['model_state_dict'])

```

*Source:* [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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:

```python
model.eval().to(device)

```

*Source:* [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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

```python

# 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

```python

# 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 `Transformer` class from [`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py) before loading weights to maintain compatibility
- **Device-agnostic loading**: Use `map_location` in `torch.load` to 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.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py) and loading in [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/generate_text.py) with architecture defined in [`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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.