# How to Save and Load Trained Transformer Models in PyTorch

> Easily save and load trained transformer models in PyTorch. Learn to store model and optimizer state_dicts using torch.save and torch.load for seamless training resumption and inference.

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

---

**Save transformer models in PyTorch by storing the model and optimizer `state_dict` dictionaries in a checkpoint file using `torch.save`, then restore them with `torch.load` and `load_state_dict` to resume training or run inference.**

The repository `FareedKhan-dev/train-llm-from-scratch` demonstrates a complete workflow for persisting GPT-style language models. By serializing only the state dictionaries rather than the full model objects, you create portable checkpoints that work across different environments and PyTorch versions.

## Understanding the Checkpoint Structure

The checkpoint system in `train-llm-from-scratch` uses a dictionary-based approach that separates model weights from training state. According to the source code in [`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py), each checkpoint contains the following keys:

- `model_state_dict`: The learnable parameters from `model.state_dict()`
- `optimizer_state_dict`: The optimizer's internal state including momentum buffers
- `losses`: Training history for plotting
- `train_loss` and `dev_loss`: Final loss values for the epoch
- `steps`: Total optimization steps completed

This structure ensures you can pause training and resume with identical optimizer momentum and learning rate schedules intact.

## Saving a Trained Transformer Model

During training, the model architecture is instantiated from [`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py) and trained via [`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py). When training completes (or at regular intervals), the script serializes the complete training state to disk.

### Complete Checkpoint Dictionary

The following code from [`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py) shows the exact pattern for creating a comprehensive checkpoint:

```python
import torch
from src.models.transformer import Transformer
from config.config import default_config as cfg

# Build the model architecture

model = Transformer(
    n_head=cfg['n_head'],
    n_embed=cfg['n_embed'],
    context_length=cfg['context_length'],
    vocab_size=cfg['vocab_size'],
    N_BLOCKS=cfg['n_blocks'],
).to(cfg['device'])

# ... training loop ...

# Save checkpoint

checkpoint_path = "checkpoints/transformer_final.pt"
torch.save(
    {
        "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),
    },
    checkpoint_path,
)

```

Storing both the **model state dict** and **optimizer state dict** allows you to pause and resume training without losing convergence progress. The file path `checkpoints/transformer_final.pt` follows the repository convention, though you can customize this location.

## Loading a Transformer Model for Inference

The [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/generate_text.py) file demonstrates how to restore a saved model for text generation. Loading requires three steps: instantiate the architecture with identical hyperparameters, load the weights, and set the model to evaluation mode.

### Step-by-Step Restoration Process

First, ensure you import the same configuration from [`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py) that was used during training. This guarantees the `Transformer` class receives identical `n_head`, `n_embed`, and `N_BLOCKS` values:

```python
import torch
import tiktoken
from src.models.transformer import Transformer
from config.config import default_config as cfg

def load_model(checkpoint_path: str, device: str = "cpu") -> Transformer:
    # Load the saved dictionary

    ckpt = torch.load(checkpoint_path, map_location=torch.device(device))
    
    # Re-instantiate the architecture

    model = Transformer(
        n_head=cfg['n_head'],
        n_embed=cfg['n_embed'],
        context_length=cfg['context_length'],
        vocab_size=cfg['vocab_size'],
        N_BLOCKS=cfg['n_blocks'],
    ).to(device)
    
    # Restore weights and set evaluation mode

    model.load_state_dict(ckpt["model_state_dict"])
    model.eval()
    return model

# Usage example

model = load_model("checkpoints/transformer_final.pt", device=cfg["device"])
enc = tiktoken.get_encoding("r50k_base")

prompt = "Once upon a time"
input_ids = torch.tensor([enc.encode_ordinary(prompt)], dtype=torch.long, device=cfg["device"])
generated = model.generate(input_ids, max_new_tokens=50)[0].tolist()
print(enc.decode(generated))

```

Calling **`.eval()`** is critical before inference because it disables dropout layers and ensures deterministic behavior during the autoregressive generation process handled by the `generate` method in [`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py).

## Resuming Training from a Checkpoint

To continue training from a saved checkpoint, you must restore both the model weights and the optimizer state. This preserves momentum buffers and learning rate schedules that are crucial for stable convergence:

```python

# Load checkpoint dict

ckpt = torch.load("checkpoints/transformer_final.pt")

# Restore model weights

model.load_state_dict(ckpt["model_state_dict"])

# Restore optimizer state (crucial for momentum-based optimizers)

optimizer.load_state_dict(ckpt["optimizer_state_dict"])

# Resume training

model.train()  # Ensure training mode is active

```

The `Transformer` class in [`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py) inherits from `nn.Module`, making it fully compatible with PyTorch's standard serialization mechanisms. Because the architecture code remains unchanged between save and load operations, the `state_dict` keys match perfectly without requiring custom deserialization logic.

## Summary

- **Checkpoint format**: Store a dictionary containing `model_state_dict`, `optimizer_state_dict`, and training metadata using `torch.save` in [`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py).
- **Architecture consistency**: Always instantiate the `Transformer` class with the same configuration from [`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py) before loading weights.
- **Inference preparation**: Call `model.eval()` after `load_state_dict()` when loading in [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/generate_text.py) to disable dropout and ensure deterministic generation.
- **Training resumption**: Restore both model and optimizer states to maintain training momentum and learning rate schedule continuity.
- **Portability**: Because only state dictionaries are stored, checkpoints work across different machines and PyTorch versions as long as the model code remains compatible.

## Frequently Asked Questions

### What is the difference between saving the entire model versus just the state_dict?

Saving the entire model with `torch.save(model, path)` pickles the complete Python object, including the class definition paths. This often breaks when moving between directories or Python versions. The `train-llm-from-scratch` repository uses `model.state_dict()` instead, which saves only the learnable tensors as ordered dictionaries. This approach requires you to recreate the model architecture with the same configuration before loading, but ensures portability and avoids pickle compatibility issues.

### Why does loading a model require instantiating the Transformer class first?

PyTorch's `load_state_dict()` method maps saved tensors to existing parameters in an instantiated model. The `Transformer` class definition in [`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py) creates the parameter structure (embeddings, attention blocks, and the language model head) that the saved weights populate. Without this existing structure, PyTorch cannot determine where to place the saved tensors. Always initialize the model with identical `n_head`, `n_embed`, and `N_BLOCKS` values from [`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py) before calling `load_state_dict()`.

### How do I ensure my checkpoint is compatible with different devices?

Use the `map_location` parameter in `torch.load()` to handle device mismatches. When loading on CPU-only machines or different GPU types, specify `map_location=torch.device(device)` as shown in [`scripts/generate_text.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/generate_text.py). This automatically moves tensors to the target device during the loading process, preventing CUDA out-of-memory errors or device mismatch runtime errors.

### Can I continue training on a different machine using these checkpoints?

Yes, provided you have the same source code for [`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py) and [`config/config.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/config/config.py). The checkpoint contains only tensor data, not the model architecture code. Copy the repository code to the new machine, ensure the configuration matches the original training environment, and load both the model and optimizer state dictionaries. The optimizer will retain its momentum buffers, allowing seamless training continuation without loss of convergence stability.