How to Save and Load Trained Transformer Models in PyTorch

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, 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 and trained via 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 shows the exact pattern for creating a comprehensive checkpoint:

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 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 that was used during training. This guarantees the Transformer class receives identical n_head, n_embed, and N_BLOCKS values:

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.

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:


# 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 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.
  • Architecture consistency: Always instantiate the Transformer class with the same configuration from config/config.py before loading weights.
  • Inference preparation: Call model.eval() after load_state_dict() when loading in 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 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 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. 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 and 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →