Skip to main content

Gradient Checkpointing in Deep Learning: How It Reduces GPU Memory Usage

Calculating read time…

Imagine you are drawing a very detailed picture — a thousand tiny steps to get from a blank page to a masterpiece. Normally, you take a photo of every single step, just in case you need to go back. That means keeping 1,000 photos in your pocket at once. Your pockets overflow!

Now imagine a smarter strategy: only keep photos at every 30th step. If you need to undo something, you go back to the nearest photo and redo just those 30 steps. A little more work — but your pockets stay empty and you can keep drawing!



That's Gradient Checkpointing in a nutshell. Instead of storing every intermediate result in your GPU's memory during training, you only keep a few key "snapshots" and recompute the rest when needed. More thinking — but far less memory. And in the world of training billion-parameter AI models, that trade-off is everything.

Gradient checkpointing is no longer optional — it is a standard, default technique used every time engineers fine-tune large language models like Llama 3, Mistral, or Falcon. If you don't know it, you'll hit an Out of Memory error and be stuck. If you do know it, you can train models that previously seemed impossible on your hardware.

📋 What You Will Learn

  • Why GPU memory is the #1 bottleneck in modern AI training
  • What happens during Forward Pass and Backward Pass (the full picture)
  • What "activations" are and why they eat so much memory
  • How Gradient Checkpointing works — step by step with diagrams
  • The exact memory vs speed trade-off (with real numbers)
  • Implementation in PyTorch — torch.utils.checkpoint
  • Implementation in Hugging Face Transformers — one line of code!
  • Implementation with DeepSpeed + FSDP for multi-GPU training
  • Selective checkpointing — the advanced technique
  • Training checkpoints (saving model state) vs gradient checkpointing — clearing the confusion
  • Full MLOps pipeline integrating gradient checkpointing with fault tolerance
  • Best practices, DOs and DON'Ts for production

1. 😱 The Memory Crisis — Why GPU Memory Is Everything

When you train a neural network, your GPU (the powerful processor used for AI) needs to hold a lot of data in its memory at once. GPU memory (called VRAM) is like a very fast but very small desk. Everything you're actively working on must fit on that desk. If it doesn't fit — crash! CUDA Out of Memory Error.

🧒 The desk analogy:

 YOUR GPU MEMORY (like a desk):

 ┌────────────────────────────────────────────────────────────────┐
 │  Model Weights  │  Gradients  │  Optimizer States  │ Activations│
 │   (the brain)   │  (learning  │  (Adam's memory)   │ (temporary │
 │                 │   signals)  │                    │  math work) │
 │  ~2-4 GB for    │  ~2-4 GB    │   ~4-8 GB          │  ~20-60 GB! │
 │  7B param model │             │                    │    😱       │
 └────────────────────────────────────────────────────────────────┘

 Total needed for a 7B param model at batch size 8: ~35-50 GB!
 A typical GPU (RTX 4090): 24 GB of VRAM.
 Problem. 😅

See that "Activations" column? It's enormous — often bigger than everything else combined! Activations are the hidden calculations stored during the forward pass. They are needed later for the backward pass (gradient calculation). This is exactly what gradient checkpointing attacks.

A 2024 research study found that for a GPT-like model with 100 billion parameters, activations alone require around 60 GB of GPU memory — even at a moderate batch size! That's more than two entire RTX 4090 cards, just for temporary math results.

2. 🔄 The Training Loop — Forward Pass + Backward Pass

To understand gradient checkpointing, you first need to understand what happens during neural network training. Every training step has two phases: Forward Pass and Backward Pass.

The Forward Pass — "Make a Prediction"

Input data flows through the neural network, layer by layer, from left to right. Each layer performs some math and produces an activation — the result of that layer's computation.

🧒 Think of it like a factory assembly line. Raw materials (input data) go in at one end. At each station (layer), a worker does some work and passes the result forward. Each worker's output is saved temporarily — it's needed to "undo" the work later.

 FORWARD PASS (making a prediction):

 Input Text                Layer 1         Layer 2         Layer 3        Output
 "Will it rain?" ──────►  [Embed] ──────► [Attention] ──► [MLP] ──────► [Prediction]
                            │               │               │
                         Save A1         Save A2         Save A3
                         (activation)    (activation)    (activation)

 A1, A2, A3 are all stored in GPU memory!
 For a 100-layer model: 100 activations stored simultaneously.
 That's HUGE.

The Backward Pass — "Learn from Mistakes"

After the forward pass gives a prediction, we compare it to the correct answer. The difference is called the loss (how wrong we were). Now we run the backward pass — flowing the error backward through every layer to compute gradients (instructions for improving each layer's weights).

🧒 Think of it like a quality inspector walking backward through the factory. They see the final product is wrong. They walk back to Layer 3 and say: "Worker at Station 3, here's what you should have done differently." Then Layer 2, then Layer 1. To do this correctly, the inspector needs to see what each worker's output was — that's the activation!

 BACKWARD PASS (learning from mistakes):

 Output ◄───────── Layer 3 needs A2 ◄──────── Layer 2 needs A1 ◄──── Input
 [Loss]             to compute its               to compute its
                    gradient                     gradient

 KEY INSIGHT:
 To compute the gradient for Layer 3, you need A2 (the output of Layer 2).
 To compute the gradient for Layer 2, you need A1 (the output of Layer 1).
 This is WHY activations must be saved during the forward pass!

3. 🧮 Three Ways to Handle Activation Memory

There are three possible strategies for managing activation memory. Gradient Checkpointing is the clever middle option.

 STRATEGY 1: SAVE EVERYTHING (Standard Training)
 ─────────────────────────────────────────────────
 Forward pass:  Save ALL activations
 Backward pass: Read them directly — fast!

 Memory:        O(n) — grows linearly with number of layers
 Speed:         Fastest possible backward pass
 Problem:       Memory explodes for deep models. 💥

 STRATEGY 2: SAVE NOTHING (Naive Memory Saving)
 ─────────────────────────────────────────────────
 Forward pass:  Save ZERO activations
 Backward pass: Recompute EVERYTHING from scratch — very slow!

 Memory:        O(1) — almost no memory needed
 Speed:         2x slower backward pass (double computation)
 Problem:       Impractically slow for real training. 🐌

 STRATEGY 3: GRADIENT CHECKPOINTING (The Smart Middle Ground) ✅
 ─────────────────────────────────────────────────────────────
 Forward pass:  Save only SELECTED activations (checkpoint nodes)
 Backward pass: Recompute only the non-saved sections

 Memory:        O(√n) — grows with SQUARE ROOT of layers!
 Speed:         ~20-30% slower than Strategy 1
 Problem:       None! This is the production standard. 🏆

🧒 The O(√n) magic:

For a model with 100 layers, Strategy 1 uses memory for 100 activations. Strategy 3 uses memory for only about 10 activations (√100 = 10). That's a 10x memory reduction for only a 20–30% speed cost! For 10,000 layers? Strategy 1 uses 10,000x memory. Strategy 3 uses 100x. Incredible!

4. 🔬 How Gradient Checkpointing Works — Step by Step

Let's trace through a 9-layer model with gradient checkpointing. We'll checkpoint every 3rd layer (so layers 3, 6, and 9 are checkpoints).

The Forward Pass with Checkpointing

 FORWARD PASS WITH CHECKPOINTING (9-layer model):

 Layer: 1    2    3    4    5    6    7    8    9
        │    │    │    │    │    │    │    │    │
        ▼    ▼    ▼    ▼    ▼    ▼    ▼    ▼    ▼
       A1   A2  [A3] A4   A5  [A6] A7   A8  [A9]
        ✗    ✗    ✓   ✗    ✗    ✓   ✗    ✗    ✓

 ✓ = SAVED in GPU memory (checkpoint node)
 ✗ = DISCARDED immediately (frees memory!)

 After forward pass, only A3, A6, A9 are in memory.
 A1, A2, A4, A5, A7, A8 have been thrown away!

 Memory used: 3 activations instead of 9 → 67% less memory! 🎉

The Backward Pass — Recomputing What Was Discarded

 BACKWARD PASS WITH CHECKPOINTING (working backward from Layer 9):

 STEP 1: Compute gradient for Layer 9
   → Needs A8. A8 was discarded!
   → Recompute from checkpoint A6: run Layer 7 → A7, run Layer 8 → A8
   → Now compute gradient for Layer 9 ✅
   → Discard A7, A8 again (no longer needed)

 STEP 2: Compute gradient for Layer 8
   → Needs A7. A7 was discarded again!
   → Recompute from checkpoint A6: run Layer 7 → A7
   → Compute gradient for Layer 8 ✅
   → Discard A7

 STEP 3: Compute gradient for Layer 7
   → Needs A6. A6 is a checkpoint → already in memory! ✅
   → No recomputation needed here.

 STEP 4: Compute gradient for Layer 6
   → Needs A5. Discarded.
   → Recompute from checkpoint A3: run Layer 4 → A4, run Layer 5 → A5
   → Compute gradient for Layer 6 ✅

 ... and so on back to Layer 1.

 Each section between checkpoints is recomputed at most ONCE.
 Total extra work = one additional forward pass through the whole network.
 That's the ~20-30% extra time cost.

The brilliant insight: every non-checkpoint activation is recomputed at most once. So the extra computation cost is bounded — it's equivalent to running one extra forward pass. And since forward passes are much cheaper than storing everything in memory, the trade-off is almost always worth it!

5. 📊 Real Numbers — What Gradient Checkpointing Actually Saves

Let's look at real benchmarks so you can see exactly what you gain (and what you give up):

 BENCHMARK: Fine-tuning Llama 3 8B model (approximate values)
 Hardware: Single NVIDIA A100 80GB GPU
 Batch size: 4, Sequence length: 2048

 ┌─────────────────────────────┬──────────────┬──────────────┬───────────┐
 │ Configuration               │ GPU Memory   │ Throughput   │ Can Run?  │
 │                             │ Used         │ (samples/sec)│           │
 ├─────────────────────────────┼──────────────┼──────────────┼───────────┤
 │ No checkpointing            │ ~72 GB       │ Fast         │ ❌ OOM!   │
 │                             │              │              │           │
 │ Gradient checkpointing ON   │ ~28 GB       │ ~80% of max  │ ✅ Works! │
 │                             │              │              │           │
 │ GC + mixed precision (fp16) │ ~18 GB       │ ~90% of max  │ ✅ Works! │
 │                             │              │              │           │
 │ GC + fp16 + LoRA (r=64)     │ ~12 GB       │ ~95% of max  │ ✅ Works! │
 └─────────────────────────────┴──────────────┴──────────────┴───────────┘

 KEY INSIGHT: Without gradient checkpointing, this model is IMPOSSIBLE to train
 on a single A100 80GB. With it, you use only 28GB. 60% memory saved!

 Activation memory reduction:   ~60-80% (typical range)
 Training speed slowdown:        ~20-30% (typical range)
 First published in: "Training Deep Nets with Sublinear Memory Cost" (Chen et al., 2016)
✅ The Golden Rule of Gradient Checkpointing:
Saving 60% memory for a 20-30% speed penalty is almost always worth it. GPU memory is the hard ceiling — you either fit in memory or you don't. A training job that's 25% slower but actually runs is infinitely better than one that crashes instantly!

6. 🚨 IMPORTANT: Two Very Different Things Called "Checkpointing"

Before diving into code, there's a critical confusion that trips up every beginner. There are two completely different things called "checkpointing" in ML:

 THING 1: GRADIENT CHECKPOINTING (= Activation Checkpointing)
 ─────────────────────────────────────────────────────────────
 WHAT:     A memory optimisation technique during training
 PURPOSE:  Save GPU memory by discarding activations and recomputing them
 WHEN:     Happens every single forward-backward pass
 SAVES:    Nothing to disk! Works entirely in GPU/CPU RAM
 AFFECTS:  Only training — has no effect on inference

 THING 2: TRAINING CHECKPOINT (= Model Checkpoint)
 ─────────────────────────────────────────────────────────────
 WHAT:     Saving the model's current state to disk
 PURPOSE:  Fault tolerance — resume training if it crashes
 WHEN:     Every N steps (e.g., every 500 steps)
 SAVES:    Files to disk: model weights, optimizer state, step number
 AFFECTS:  Resumability and reproducibility

 ANALOGY:
 Gradient Checkpointing = saving your RAM by strategically forgetting
                          intermediate math (and recomputing if needed)
 Training Checkpoint    = hitting Ctrl+S to save your document to disk
                          so you don't lose work if the computer crashes
⚠️ They are completely independent!
You can use gradient checkpointing (memory trick) without saving training checkpoints (disk saves), and vice versa. In production you almost always use BOTH together. This guide covers both in detail!

7. 💻 Implementation — PyTorch Native (torch.utils.checkpoint)

PyTorch provides gradient checkpointing directly out of the box with just a few lines of code. You don't need any extra libraries — it's built right in!

Install libraries (if needed)

📝 What the command below does:
Installs PyTorch (the deep learning framework) with CUDA support (for GPU training). torchvision and torchaudio are standard companions. Visit pytorch.org for the exact command for your operating system and CUDA version!
# Install PyTorch with CUDA support
# Visit pytorch.org for the exact command for your OS and CUDA version
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121
pip install transformers accelerate datasets

Method 1: Using torch.utils.checkpoint.checkpoint() on Custom Models

📝 What the code below does:
We build a simple neural network with 6 layers. Without gradient checkpointing, all 6 layers' intermediate results (activations) would be stored in GPU memory during the forward pass. With gradient checkpointing, we wrap specific groups of layers using checkpoint(). PyTorch will discard those layers' intermediate results during the forward pass and recompute them during the backward pass — saving memory at the cost of slightly more computation. We also show how to measure the memory usage before and after to see the real difference!
# gradient_checkpoint_demo.py
# Demonstrates gradient checkpointing on a custom PyTorch model
# and measures the memory savings

import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint

# ── Helper function to measure GPU memory usage ───────────────────────────────
def get_gpu_memory_mb():
    """Returns the current GPU memory allocated in MB."""
    if torch.cuda.is_available():
        return torch.cuda.memory_allocated() / 1024**2
    return 0

def reset_gpu_memory():
    """Clear GPU memory cache for a fresh measurement."""
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.reset_peak_memory_stats()

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")

# ── Build a simple but deep neural network ─────────────────────────────────────
# This simulates a small version of a transformer's MLP block
# Each "block" is a Linear layer + ReLU activation
class MLPBlock(nn.Module):
    """One block of our neural network: Linear → ReLU → Linear → ReLU"""
    def __init__(self, hidden_size=2048):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(hidden_size, hidden_size),
            nn.ReLU(),
            nn.Linear(hidden_size, hidden_size),
            nn.ReLU(),
        )

    def forward(self, x):
        return self.net(x)

# ────────────────────────────────────────────────────────────────────────────────
# VERSION 1: Standard training (NO gradient checkpointing)
# All layer activations are stored in GPU memory throughout the forward pass
# ────────────────────────────────────────────────────────────────────────────────
class StandardModel(nn.Module):
    """Standard model: stores ALL activations during forward pass."""
    def __init__(self, n_blocks=12, hidden_size=2048):
        super().__init__()
        self.blocks = nn.ModuleList([MLPBlock(hidden_size) for _ in range(n_blocks)])

    def forward(self, x):
        for block in self.blocks:
            x = block(x)   # Activation stored for every block → lots of memory!
        return x

# ────────────────────────────────────────────────────────────────────────────────
# VERSION 2: Gradient checkpointing model
# We wrap each block with checkpoint() → activations are discarded and recomputed
# ────────────────────────────────────────────────────────────────────────────────
class CheckpointedModel(nn.Module):
    """Checkpointed model: discards activations during forward, recomputes in backward."""
    def __init__(self, n_blocks=12, hidden_size=2048):
        super().__init__()
        self.blocks = nn.ModuleList([MLPBlock(hidden_size) for _ in range(n_blocks)])

    def forward(self, x):
        for block in self.blocks:
            # checkpoint() wraps the block.forward function
            # use_reentrant=False: use the modern, safer non-reentrant API
            # (recommended for all new code — the old reentrant version has known bugs)
            x = checkpoint(block, x, use_reentrant=False)
        return x

# ── Measure memory for STANDARD model ─────────────────────────────────────────
print("\n" + "═"*60)
print("Comparing memory usage: Standard vs Checkpointed")
print("═"*60)

reset_gpu_memory()
standard_model = StandardModel(n_blocks=12, hidden_size=2048).to(device)
optimizer_std  = torch.optim.Adam(standard_model.parameters(), lr=1e-4)

# Create fake input: batch of 16 samples, each with 2048 features
x_input = torch.randn(16, 2048, requires_grad=True).to(device)

# Standard forward + backward pass
standard_model.train()
output_std = standard_model(x_input)
loss_std   = output_std.mean()
loss_std.backward()

std_memory = get_gpu_memory_mb()
print(f"\n📊 Standard Model (no checkpointing):")
print(f"   Peak GPU Memory: {torch.cuda.max_memory_allocated()/1024**2:.1f} MB")

# ── Measure memory for CHECKPOINTED model ─────────────────────────────────────
reset_gpu_memory()
checkpointed_model = CheckpointedModel(n_blocks=12, hidden_size=2048).to(device)
optimizer_chk      = torch.optim.Adam(checkpointed_model.parameters(), lr=1e-4)

checkpointed_model.train()
output_chk = checkpointed_model(x_input)
loss_chk   = output_chk.mean()
loss_chk.backward()

print(f"\n📊 Checkpointed Model (gradient checkpointing ON):")
print(f"   Peak GPU Memory: {torch.cuda.max_memory_allocated()/1024**2:.1f} MB")
print(f"\n✅ Memory saved: see the difference above!")
print(f"   Typical saving: 30-70% depending on model architecture.")

Key lines explained:

  • from torch.utils.checkpoint import checkpoint → Import the built-in PyTorch checkpointing function. No extra install needed!
  • checkpoint(block, x, use_reentrant=False) → The magic line. Wraps one forward function call with checkpointing. use_reentrant=False is the modern, safer version — always use this for new code.
  • torch.cuda.max_memory_allocated() → PyTorch's built-in tool to measure peak GPU memory usage during training.
  • torch.cuda.empty_cache() → Clears PyTorch's internal memory cache. Call before benchmarks for clean measurements.

Method 2: checkpoint_sequential() for Sequential Models

📝 What the code below does:
If your model uses nn.Sequential (layers stacked one after another in order), PyTorch provides a special helper called checkpoint_sequential(). Instead of wrapping each layer individually, you just tell it your sequential model and how many chunks to split it into. It automatically places checkpoints at equal intervals! This is the fastest way to add checkpointing to a sequential architecture.
# checkpoint_sequential_demo.py
# A faster way to add checkpointing to nn.Sequential models

import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint_sequential

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# Build a deep sequential model (like a simplified ResNet or MLP)
hidden = 1024
sequential_model = nn.Sequential(
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
    nn.Linear(hidden, hidden), nn.ReLU(),
).to(device)

sequential_model.train()
x = torch.randn(32, hidden, requires_grad=True).to(device)

# checkpoint_sequential(model, segments, input)
# segments=4 means: split the model into 4 equal chunks
# Checkpoints are placed automatically between each chunk
# Chunk 1: Layers 1-2, Chunk 2: Layers 3-4, etc.
# Only the output of each chunk is saved — intermediate values discarded!
output = checkpoint_sequential(
    sequential_model,
    segments=4,     # Split into 4 chunks → 3 checkpoint points
    input=x,
    use_reentrant=False   # Always use False for new code
)

loss = output.mean()
loss.backward()

print("✅ checkpoint_sequential() completed successfully!")
print(f"   GPU memory: {torch.cuda.memory_allocated()/1024**2:.1f} MB")

8. 🤗 The Easiest Way — Hugging Face Transformers (One Line!)

If you are fine-tuning a Hugging Face model (like BERT, GPT-2, Llama, Mistral), enabling gradient checkpointing is literally one line of code. Hugging Face handles all the complexity for you automatically!

📝 What the code below does:
We load a pre-trained language model from Hugging Face (Llama 3.2 1B in this example — small enough to run on most GPUs). Then we call model.gradient_checkpointing_enable() — that's it! Under the hood, Hugging Face applies torch.utils.checkpoint to every Transformer block automatically. We also show the full training setup with the Hugging Face Trainer, where gradient checkpointing is a simple config option.
# hf_gradient_checkpointing.py
# Gradient checkpointing with Hugging Face Transformers — the easy way!

from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    TrainingArguments,
    Trainer,
    DataCollatorForLanguageModeling
)
from datasets import load_dataset
import torch

# ── Method A: Manual gradient checkpointing enable ────────────────────────────
# This is the direct API — use this when writing your own training loop

model_name = "meta-llama/Llama-3.2-1B"   # 1B parameter model — small but real!
# For testing without downloading, use: "gpt2" (124M params, freely available)

print(f"Loading model: {model_name}")
model = AutoModelForCausalLM.from_pretrained(
    "gpt2",             # Using GPT-2 for easy download — replace with Llama for production
    torch_dtype=torch.float16   # fp16 mixed precision saves additional memory
)

# ── THE MAGIC LINE ─────────────────────────────────────────────────────────────
# This single call applies gradient checkpointing to every Transformer block.
# No need to understand the internals — Hugging Face handles it all!
model.gradient_checkpointing_enable(
    gradient_checkpointing_kwargs={"use_reentrant": False}  # Use modern non-reentrant API
)

print(f"✅ Gradient checkpointing enabled!")
print(f"   Is it enabled? {model.is_gradient_checkpointing}")   # Should print True

# ── IMPORTANT: Disable cache when using gradient checkpointing ─────────────────
# The model's KV cache (key-value cache used for attention) is incompatible
# with gradient checkpointing. Disable it during training!
# (The cache is for inference speed — not needed during training anyway)
model.config.use_cache = False

print(f"   KV cache disabled: {not model.config.use_cache}")

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = model.to(device)

# ── Method B: Using Hugging Face Trainer (recommended for production) ──────────
# When using Trainer, set gradient_checkpointing=True in TrainingArguments
# The Trainer handles everything: enabling checkpointing, disabling cache, etc.

print("\n" + "═"*60)
print("Setting up Trainer with gradient checkpointing...")
print("═"*60)

training_args = TrainingArguments(
    output_dir                    = "./output",
    num_train_epochs              = 1,
    per_device_train_batch_size   = 4,        # Batch size per GPU
    gradient_accumulation_steps   = 4,        # Effective batch size = 4 × 4 = 16
    # ── Memory optimization settings ──────────────────────────────────────────
    gradient_checkpointing        = True,     # ← ENABLE GRADIENT CHECKPOINTING!
    gradient_checkpointing_kwargs = {"use_reentrant": False},
    fp16                          = True,     # Mixed precision training (saves more memory)
    # ── Training settings ─────────────────────────────────────────────────────
    learning_rate                 = 2e-5,
    warmup_ratio                  = 0.1,
    logging_steps                 = 10,
    save_steps                    = 500,      # Save training checkpoint every 500 steps
    save_total_limit              = 3,        # Only keep last 3 checkpoints (saves disk)
    report_to                     = "none",   # Set to "wandb" or "mlflow" for tracking
)

print("✅ TrainingArguments configured!")
print(f"   gradient_checkpointing: {training_args.gradient_checkpointing}")
print(f"   fp16 mixed precision:   {training_args.fp16}")
print(f"   effective batch size:   "
      f"{training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps}")

Key lines explained:

  • model.gradient_checkpointing_enable() → One-line activation. Hugging Face applies checkpointing to every Transformer block internally.
  • use_reentrant=False → Always pass this! The old reentrant mode has known bugs with some model architectures in distributed training.
  • model.config.use_cache = False → Critical! The attention KV cache conflicts with gradient checkpointing. Must be disabled during training.
  • gradient_accumulation_steps=4 → A companion technique — accumulates gradients over 4 small batches before updating. Simulates larger batch sizes without memory cost.
  • save_total_limit=3 → Keep only the last 3 training checkpoints on disk to save storage.

9. 🔁 Full Training Loop with Gradient Checkpointing + Training Checkpoints

Now let's put it all together — gradient checkpointing (memory) AND training checkpoints (disk saves) in one complete production-ready training loop. This is the pattern used by professional ML engineers to fine-tune large models!

📝 What the code below does:
This is a complete, production-ready training loop combining both types of checkpointing. Gradient checkpointing (memory technique) is enabled on the model. Training checkpoints (disk saves) are saved every 200 steps, recording the model weights, optimizer state, and exactly where training stopped. If training crashes at step 1,847 (GPU failure, cloud preemption, power cut), you just reload the checkpoint from step 1,800 and continue — losing only 47 steps instead of everything!
# complete_training_loop.py
# Production-ready training loop with BOTH types of checkpointing

import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
from transformers import AutoModelForCausalLM, AutoTokenizer, get_linear_schedule_with_warmup
import os
import json

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# ── Configuration ──────────────────────────────────────────────────────────────
CHECKPOINT_DIR       = "./training_checkpoints"    # Where to save disk checkpoints
SAVE_EVERY_N_STEPS   = 200                         # Save to disk every N steps
GRADIENT_ACCUM_STEPS = 4                           # Accumulate gradients over N steps
LEARNING_RATE        = 2e-5
NUM_EPOCHS           = 3
os.makedirs(CHECKPOINT_DIR, exist_ok=True)

# ── Load model with gradient checkpointing enabled ────────────────────────────
print("Loading model and enabling gradient checkpointing...")
model = AutoModelForCausalLM.from_pretrained("gpt2", torch_dtype=torch.float16)
model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
model.config.use_cache = False   # Required when using gradient checkpointing!
model = model.to(device)

# ── Set up optimizer and scheduler ────────────────────────────────────────────
optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=0.01)

# Total training steps for scheduler (example: 1000 steps of fake data)
total_steps = 1000
scheduler   = get_linear_schedule_with_warmup(
    optimizer, num_warmup_steps=int(total_steps * 0.1), num_training_steps=total_steps
)

# ── Function to save a training checkpoint ────────────────────────────────────
def save_training_checkpoint(step, epoch, model, optimizer, scheduler, loss):
    """
    Saves everything needed to resume training exactly where we stopped.
    This is the "Ctrl+S" for your training run!
    Think of it like a save point in a video game.
    """
    checkpoint_path = os.path.join(CHECKPOINT_DIR, f"checkpoint_step_{step}.pt")

    torch.save({
        "step":                    step,             # Exactly where we are
        "epoch":                   epoch,            # Which epoch
        "model_state_dict":        model.state_dict(),     # All model weights
        "optimizer_state_dict":    optimizer.state_dict(), # Optimizer memory (Adam needs this!)
        "scheduler_state_dict":    scheduler.state_dict(), # Learning rate schedule state
        "loss":                    loss,             # Loss value at this checkpoint
    }, checkpoint_path)

    # Save metadata as human-readable JSON for quick inspection
    metadata = {"step": step, "epoch": epoch, "loss": float(loss),
                "checkpoint_file": checkpoint_path}
    with open(os.path.join(CHECKPOINT_DIR, f"checkpoint_step_{step}_meta.json"), "w") as f:
        json.dump(metadata, f, indent=2)

    print(f"  💾 Training checkpoint saved: step {step} → {checkpoint_path}")

    # Clean up old checkpoints — keep only the last 3 to save disk space
    checkpoint_files = sorted([
        f for f in os.listdir(CHECKPOINT_DIR) if f.endswith(".pt")
    ])
    while len(checkpoint_files) > 3:
        old_checkpoint = os.path.join(CHECKPOINT_DIR, checkpoint_files.pop(0))
        os.remove(old_checkpoint)
        print(f"  🗑️  Removed old checkpoint: {old_checkpoint}")

# ── Function to resume from a training checkpoint ─────────────────────────────
def load_latest_checkpoint(model, optimizer, scheduler):
    """
    Finds the most recent checkpoint and resumes training from there.
    Call this at the start of training to automatically continue if a crash happened.
    """
    checkpoint_files = sorted([
        f for f in os.listdir(CHECKPOINT_DIR) if f.endswith(".pt")
    ])
    if not checkpoint_files:
        print("  No existing checkpoints found — starting fresh training.")
        return 0, 0   # start_step=0, start_epoch=0

    latest_path = os.path.join(CHECKPOINT_DIR, checkpoint_files[-1])
    print(f"  📂 Resuming from checkpoint: {latest_path}")

    checkpoint_data = torch.load(latest_path, map_location=device)
    model.load_state_dict(checkpoint_data["model_state_dict"])
    optimizer.load_state_dict(checkpoint_data["optimizer_state_dict"])
    scheduler.load_state_dict(checkpoint_data["scheduler_state_dict"])

    resume_step  = checkpoint_data["step"]
    resume_epoch = checkpoint_data["epoch"]
    print(f"  ✅ Resumed from step {resume_step}, epoch {resume_epoch}")
    return resume_step, resume_epoch

# ── Check if we should resume from a previous run ─────────────────────────────
start_step, start_epoch = load_latest_checkpoint(model, optimizer, scheduler)

# ── Main training loop ─────────────────────────────────────────────────────────
print(f"\n🚀 Starting training from step {start_step}...")

model.train()
global_step = start_step

# Simulate a training dataset with fake tokens
# In real training: replace with your DataLoader
import numpy as np
fake_batches_per_epoch = 250

for epoch in range(start_epoch, NUM_EPOCHS):
    print(f"\n{'═'*50}")
    print(f"Epoch {epoch + 1}/{NUM_EPOCHS}")
    print(f"{'═'*50}")

    total_loss    = 0.0
    optimizer.zero_grad()   # Clear gradients at the start

    for batch_idx in range(fake_batches_per_epoch):

        # ── Create fake batch (replace with real DataLoader in production) ─────
        input_ids = torch.randint(0, 50257, (2, 128)).to(device)  # batch=2, seq_len=128
        labels    = input_ids.clone()

        # ── Forward pass (gradient checkpointing happens automatically here!) ──
        # Because we called gradient_checkpointing_enable(), PyTorch automatically
        # discards intermediate activations and recomputes them in the backward pass.
        # You don't need to change ANYTHING in this line!
        outputs = model(input_ids=input_ids, labels=labels)
        loss    = outputs.loss

        # Normalize loss by accumulation steps so the scale stays correct
        loss = loss / GRADIENT_ACCUM_STEPS

        # ── Backward pass ──────────────────────────────────────────────────────
        loss.backward()   # Gradient checkpointing recomputation happens HERE

        total_loss += loss.item()

        # ── Optimizer step (only after accumulating enough gradients) ──────────
        if (batch_idx + 1) % GRADIENT_ACCUM_STEPS == 0:
            # Gradient clipping: prevent gradients from exploding
            # max_norm=1.0 is a standard setting for language model training
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

            optimizer.step()    # Update model weights
            scheduler.step()    # Update learning rate
            optimizer.zero_grad()   # Clear gradients for next accumulation

            global_step += 1

            # ── Save training checkpoint every N steps ─────────────────────────
            if global_step % SAVE_EVERY_N_STEPS == 0:
                current_loss = total_loss / SAVE_EVERY_N_STEPS
                save_training_checkpoint(
                    step=global_step, epoch=epoch,
                    model=model, optimizer=optimizer, scheduler=scheduler,
                    loss=current_loss
                )
                total_loss = 0.0   # Reset running loss

            # ── Progress logging ───────────────────────────────────────────────
            if global_step % 50 == 0:
                lr_now = scheduler.get_last_lr()[0]
                mem_mb = torch.cuda.memory_allocated() / 1024**2 if torch.cuda.is_available() else 0
                print(f"  Step {global_step:5d} | "
                      f"Loss: {loss.item() * GRADIENT_ACCUM_STEPS:.4f} | "
                      f"LR: {lr_now:.2e} | "
                      f"GPU Mem: {mem_mb:.0f} MB")

print(f"\n✅ Training complete! Final step: {global_step}")
print(f"   All checkpoints saved in: {CHECKPOINT_DIR}")

10. ⚡ Advanced — Gradient Checkpointing with DeepSpeed + FSDP

When training very large models (7B+ parameters) across multiple GPUs, you combine gradient checkpointing with distributed training frameworks. The two most important ones are DeepSpeed and FSDP.

 WHEN TO USE WHICH FRAMEWORK (Guide):

 ┌──────────────────┬──────────────────────────────────────────────────────┐
 │ Situation        │ Recommended Approach                                 │
 ├──────────────────┼──────────────────────────────────────────────────────┤
 │ Single GPU       │ gradient_checkpointing_enable() + fp16               │
 │                  │ (This guide, Hugging Face Trainer)                   │
 ├──────────────────┼──────────────────────────────────────────────────────┤
 │ Multiple GPUs,   │ PyTorch FSDP + activation checkpointing              │
 │ model fits in    │ (Cleaner, native PyTorch, good docs)                 │
 │ aggregate memory │                                                      │
 ├──────────────────┼──────────────────────────────────────────────────────┤
 │ Huge models      │ DeepSpeed ZeRO-3 + activation checkpointing          │
 │ 70B+ params      │ (CPU offloading, NVMe offloading, most powerful)     │
 │ or NVMe offload  │                                                      │
 ├──────────────────┼──────────────────────────────────────────────────────┤
 │ Fine-tuning with │ DDP or FSDP + LoRA + gradient checkpointing          │
 │ LoRA/QLoRA       │ (Simpler setup, no need for DeepSpeed)               │
 └──────────────────┴──────────────────────────────────────────────────────┘

DeepSpeed Configuration with Gradient Checkpointing

📝 What the file below does:
This is a DeepSpeed configuration file (JSON format). It tells DeepSpeed all the memory optimisation settings for a multi-GPU training run. The "activation_checkpointing" section enables gradient checkpointing at the DeepSpeed level — it works alongside PyTorch's gradient checkpointing for maximum memory savings. ZeRO Stage 2 shards optimizer states and gradients across GPUs (saves lots of VRAM). partition_activations splits checkpoint activations across GPUs too — even more memory savings!
{
  "comment": "DeepSpeed config with gradient checkpointing + ZeRO Stage 2",

  "zero_optimization": {
    "stage": 2,
    "allgather_partitions": true,
    "allgather_bucket_size": 2e8,
    "overlap_comm": true,
    "reduce_scatter": true,
    "reduce_bucket_size": 2e8,
    "contiguous_gradients": true
  },

  "activation_checkpointing": {
    "partition_activations": true,
    "cpu_checkpointing": false,
    "contiguous_memory_optimization": true,
    "number_checkpoints": null,
    "synchronize_checkpoint_boundary": false,
    "profile": false
  },

  "fp16": {
    "enabled": true,
    "loss_scale": 0,
    "loss_scale_window": 1000,
    "initial_scale_power": 16,
    "hysteresis": 2,
    "min_loss_scale": 1
  },

  "optimizer": {
    "type": "AdamW",
    "params": {
      "lr": 2e-5,
      "betas": [0.9, 0.999],
      "eps": 1e-8,
      "weight_decay": 0.01
    }
  },

  "gradient_clipping": 1.0,
  "train_batch_size": 16,
  "train_micro_batch_size_per_gpu": 2,
  "gradient_accumulation_steps": 8,
  "steps_per_print": 50,
  "wall_clock_breakdown": false
}
📝 What the code below does:
This Python script uses Hugging Face Accelerate with the DeepSpeed config above. Accelerate is a wrapper library that makes it easy to use DeepSpeed, FSDP, or single-GPU training with the same training code — just change a config file! We enable gradient checkpointing on the model AND in the DeepSpeed config. This launches a multi-GPU training run from the terminal.
# deepspeed_training.py
# Multi-GPU fine-tuning with DeepSpeed + gradient checkpointing
# Run with: accelerate launch --config_file accelerate_config.yaml deepspeed_training.py
# OR:       deepspeed --num_gpus=2 deepspeed_training.py

from accelerate import Accelerator
from accelerate.utils import DeepSpeedPlugin
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch

# ── Set up Accelerate with DeepSpeed ──────────────────────────────────────────
# Accelerate automatically detects whether to use DeepSpeed based on your config
accelerator = Accelerator()

print(f"Number of GPUs: {accelerator.num_processes}")
print(f"Distributed type: {accelerator.distributed_type}")

# ── Load model ────────────────────────────────────────────────────────────────
model = AutoModelForCausalLM.from_pretrained("gpt2")

# ── Enable gradient checkpointing ─────────────────────────────────────────────
# Call this BEFORE wrapping with accelerator.prepare()
model.gradient_checkpointing_enable(
    gradient_checkpointing_kwargs={"use_reentrant": False}
)
model.config.use_cache = False   # Always disable KV cache during training!

print("✅ Gradient checkpointing enabled on model")

# ── Prepare with Accelerate (handles DeepSpeed sharding automatically) ─────────
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5)
model, optimizer = accelerator.prepare(model, optimizer)

print("✅ Model and optimizer prepared with Accelerate")
print("   (DeepSpeed ZeRO sharding applied automatically!)")

# ── Your training loop goes here ───────────────────────────────────────────────
# The training loop is identical to single-GPU training!
# Accelerate handles all the distributed communication under the hood.

model.train()
for step in range(10):  # Short demo loop
    fake_input = torch.randint(0, 50257, (2, 64))
    outputs = model(input_ids=fake_input, labels=fake_input)
    loss = outputs.loss

    accelerator.backward(loss)   # Use accelerator.backward() instead of loss.backward()!

    if (step + 1) % 2 == 0:
        optimizer.step()
        optimizer.zero_grad()
        print(f"  Step {step+1}: Loss = {loss.item():.4f}")

print("\n✅ DeepSpeed multi-GPU training demo complete!")

11. 🎯 Advanced: Selective Checkpointing (The Expert Technique)

Standard gradient checkpointing applies checkpointing to every Transformer block equally. But not all layers cost the same to recompute! Selective checkpointing is smarter: apply checkpointing only to the layers where it saves the most memory with the least recomputation cost.

🧒 Analogy: Instead of taking a photo at every 3rd step of painting, you're smarter now. You only skip photos for the simple strokes (quick to redo) and always keep photos for the complex ones (expensive to redo). More efficient use of your "pocket space"!

 SELECTIVE CHECKPOINTING STRATEGY:

 Transformer Block Types and their costs:

 ┌─────────────────────┬────────────────┬─────────────────────────────┐
 │ Layer Type          │ Memory Cost    │ Recompute Cost   │ Checkpoint?│
 ├─────────────────────┼────────────────┼──────────────────┼────────────┤
 │ Self-Attention      │ VERY HIGH 😱   │ Medium           │ YES ✅     │
 │ (KV activations)    │ scales with    │                  │            │
 │                     │ seq_len²       │                  │            │
 ├─────────────────────┼────────────────┼──────────────────┼────────────┤
 │ MLP / FFN Layer     │ HIGH           │ Low              │ YES ✅     │
 │                     │ 4x hidden_dim  │                  │            │
 ├─────────────────────┼────────────────┼──────────────────┼────────────┤
 │ LayerNorm           │ Very Low       │ Very Low         │ NO ❌      │
 │                     │                │                  │ (not worth)│
 ├─────────────────────┼────────────────┼──────────────────┼────────────┤
 │ Embedding Layer     │ Medium         │ Very Low         │ NO ❌      │
 │                     │                │                  │            │
 └─────────────────────┴────────────────┴──────────────────┴────────────┘

 Best practice: Checkpoint every Transformer block (Attention + MLP together)
 but skip lightweight operations like LayerNorm and activation functions.
📝 What the code below does:
This shows how to apply selective gradient checkpointing using PyTorch FSDP (Fully Sharded Data Parallel). We use apply_activation_checkpointing() with a custom check_fn function that tells PyTorch exactly which types of modules should be checkpointed. Only TransformerBlock instances get checkpointed — not embedding layers, not LayerNorms. This gives better memory savings with lower compute overhead than checkpointing everything!
# selective_checkpointing.py
# Apply gradient checkpointing selectively to specific layer types

import torch
import torch.nn as nn
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
    checkpoint_wrapper,
    CheckpointImpl,
    apply_activation_checkpointing,
)

# ── Define a simple Transformer-style model ────────────────────────────────────
class TransformerBlock(nn.Module):
    """One Transformer block: Attention + MLP. This is the main memory consumer."""
    def __init__(self, hidden_size=768):
        super().__init__()
        self.attention = nn.MultiheadAttention(hidden_size, num_heads=12, batch_first=True)
        self.norm1     = nn.LayerNorm(hidden_size)
        self.mlp       = nn.Sequential(
            nn.Linear(hidden_size, hidden_size * 4),
            nn.GELU(),
            nn.Linear(hidden_size * 4, hidden_size),
        )
        self.norm2 = nn.LayerNorm(hidden_size)

    def forward(self, x):
        # Attention sub-block
        attn_out, _ = self.attention(x, x, x)
        x = self.norm1(x + attn_out)
        # MLP sub-block
        x = self.norm2(x + self.mlp(x))
        return x

class EmbeddingLayer(nn.Module):
    """The embedding layer — converts token IDs to vectors. Small memory cost."""
    def __init__(self, vocab_size=50257, hidden_size=768):
        super().__init__()
        self.embed = nn.Embedding(vocab_size, hidden_size)

    def forward(self, x):
        return self.embed(x)

class SimpleTransformer(nn.Module):
    """A simple Transformer with embedding + 6 transformer blocks."""
    def __init__(self, vocab_size=50257, hidden_size=768, n_layers=6):
        super().__init__()
        self.embedding = EmbeddingLayer(vocab_size, hidden_size)
        self.blocks    = nn.ModuleList([TransformerBlock(hidden_size) for _ in range(n_layers)])
        self.head      = nn.Linear(hidden_size, vocab_size)

    def forward(self, x):
        x = self.embedding(x)
        for block in self.blocks:
            x = block(x)
        return self.head(x)

# ── Create model ──────────────────────────────────────────────────────────────
model = SimpleTransformer(n_layers=6)

# ── Apply SELECTIVE gradient checkpointing ────────────────────────────────────
# check_fn tells PyTorch which modules to checkpoint.
# We ONLY checkpoint TransformerBlock instances — not embedding, not head.
# Why? TransformerBlocks are expensive in memory. Embedding/head are cheap.
check_fn = lambda submodule: isinstance(submodule, TransformerBlock)

apply_activation_checkpointing(
    model,
    checkpoint_wrapper_fn = lambda module: checkpoint_wrapper(
        module,
        checkpoint_impl=CheckpointImpl.NO_REENTRANT   # Modern non-reentrant mode
    ),
    check_fn = check_fn   # Only apply to TransformerBlock instances
)

print("✅ Selective gradient checkpointing applied!")
print(f"   Checkpointed: TransformerBlock layers (6 of them)")
print(f"   Skipped:      EmbeddingLayer, Linear head, LayerNorm")
print(f"   Result: Maximum memory saving, minimum recompute overhead!")

# Verify: do a forward + backward pass to confirm it works
model.train()
x_test = torch.randint(0, 50257, (2, 32))   # batch=2, seq_len=32
output = model(x_test)
loss   = output.mean()
loss.backward()
print(f"\n✅ Forward + backward pass successful with selective checkpointing!")

12. 🔍 Monitoring GPU Memory in MLOps

In production MLOps, you need to measure your memory usage, not just hope it works. Here's how to profile GPU memory during training:

📝 What the code below does:
These are the standard tools for monitoring GPU memory during training. We show three methods: (1) PyTorch's built-in memory tracker — simplest, no setup needed, (2) The NVIDIA CLI tool nvidia-smi for real-time GPU monitoring from your terminal, (3) PyTorch's full memory profiler for finding exactly which operation causes memory spikes. Use these to verify that gradient checkpointing is actually reducing your memory usage!
# memory_monitoring.py
# Tools for monitoring GPU memory usage during training

import torch

# ══════════════════════════════════════════════════════════════════════════
# METHOD 1: PyTorch built-in memory functions (quick and easy)
# Use these to print memory stats at any point during training
# ══════════════════════════════════════════════════════════════════════════

def print_memory_stats(label=""):
    """Print current and peak GPU memory usage. Call at any point in training."""
    if not torch.cuda.is_available():
        print("No GPU available — memory stats N/A")
        return

    allocated  = torch.cuda.memory_allocated()  / 1024**2   # MB currently used
    reserved   = torch.cuda.memory_reserved()   / 1024**2   # MB reserved (may differ from used)
    peak_alloc = torch.cuda.max_memory_allocated()/ 1024**2  # Peak usage since last reset

    print(f"\n📊 GPU Memory [{label}]:")
    print(f"   Allocated (current):  {allocated:.1f} MB")
    print(f"   Reserved (PyTorch):   {reserved:.1f} MB")
    print(f"   Peak (since reset):   {peak_alloc:.1f} MB")

def reset_memory_stats():
    """Reset peak memory tracking. Call before a section you want to benchmark."""
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.reset_peak_memory_stats()

# ── Example: compare memory before and after enabling checkpointing ────────────
import torch.nn as nn
from torch.utils.checkpoint import checkpoint

device = "cuda" if torch.cuda.is_available() else "cpu"

simple_model = nn.Sequential(
    *[nn.Sequential(nn.Linear(512, 512), nn.ReLU()) for _ in range(20)]
).to(device)
simple_model.train()
x_demo = torch.randn(8, 512, requires_grad=True).to(device)

print("\n" + "═"*50)
print("Memory Comparison: Standard vs Checkpointed")
print("═"*50)

# WITHOUT checkpointing
reset_memory_stats()
y = x_demo
for layer in simple_model:
    y = layer(y)
y.mean().backward()
print_memory_stats("STANDARD (no checkpointing)")

# WITH checkpointing
reset_memory_stats()
y = x_demo
for layer in simple_model:
    y = checkpoint(layer, y, use_reentrant=False)
y.mean().backward()
print_memory_stats("CHECKPOINTED")

# ══════════════════════════════════════════════════════════════════════════
# METHOD 2: nvidia-smi — monitor from your terminal (no code needed!)
# Run this in a SEPARATE terminal while training is running:
# ══════════════════════════════════════════════════════════════════════════

nvidia_smi_commands = """
# Monitor all GPUs every 1 second (live view):
nvidia-smi -l 1

# Show memory summary for all GPUs:
nvidia-smi --query-gpu=gpu_name,memory.used,memory.free,memory.total,utilization.gpu \
           --format=csv,noheader

# Detailed memory breakdown (fragmentation info):
nvidia-smi pmon -s mu -d 5
"""
print(f"\n💡 Monitor GPU memory from terminal:")
print(nvidia_smi_commands)

# ══════════════════════════════════════════════════════════════════════════
# METHOD 3: PyTorch Memory Profiler (find exact memory hotspots)
# Records a detailed trace of every memory allocation during training
# ══════════════════════════════════════════════════════════════════════════

print("\n" + "═"*50)
print("PyTorch Memory Profiler (advanced)")
print("═"*50)

with torch.profiler.profile(
    activities=[torch.profiler.ProfilerActivity.CPU,
                torch.profiler.ProfilerActivity.CUDA],
    profile_memory=True,      # Track memory allocations and frees
    record_shapes=True,       # Record tensor shapes for debugging
    with_stack=False          # Set True for full Python stack traces (slower)
) as prof:
    # Run a few training steps inside the profiler context
    for _ in range(3):
        x_prof = torch.randn(4, 512, requires_grad=True).to(device)
        y_prof = x_prof
        for layer in list(simple_model)[:5]:   # Profile first 5 layers
            y_prof = checkpoint(layer, y_prof, use_reentrant=False)
        y_prof.mean().backward()

# Print the operations that used the most memory
print("\nTop 10 memory-consuming operations:")
print(prof.key_averages().table(
    sort_by="self_cuda_memory_usage",
    row_limit=10
))

print("\n✅ Memory profiling complete!")
print("   Save trace with: prof.export_chrome_trace('trace.json')")
print("   Open in Chrome at: chrome://tracing  (amazing visualisation!)")

13. 🏭 Complete MLOps Pipeline with Gradient Checkpointing

Here's the full picture of how gradient checkpointing fits into a production MLOps workflow:

 COMPLETE MLOPS PIPELINE WITH GRADIENT CHECKPOINTING 
 ═══════════════════════════════════════════════════════════════════════

 PHASE 1: ENVIRONMENT SETUP
 ─────────────────────────────────────────────────────────────────────
 Install: PyTorch, Transformers, Accelerate, DeepSpeed
 Configure: accelerate_config.yaml or deepspeed_config.json
 Set GPU memory budget → decide checkpointing strategy
 Measure baseline memory (no checkpointing) → confirm OOM situation

         ↓

 PHASE 2: MODEL + MEMORY OPTIMISATION STACK
 ─────────────────────────────────────────────────────────────────────
 Layer 1: Gradient Checkpointing          → saves ~60% activation memory
 Layer 2: Mixed Precision (fp16/bf16)     → saves ~50% weight/gradient memory
 Layer 3: Gradient Accumulation           → simulate larger batch sizes
 Layer 4: LoRA / QLoRA (optional)         → reduce trainable params by 100x
 Layer 5: DeepSpeed ZeRO / FSDP          → shard across multiple GPUs

         ↓

 PHASE 3: TRAINING RUN
 ─────────────────────────────────────────────────────────────────────
 gradient_checkpointing_enable()      ← activate memory trick
 monitor memory every N steps         ← alert if approaching limit
 save training checkpoint every M steps ← fault tolerance
 log metrics to MLflow / W&B          ← experiment tracking

         ↓

 PHASE 4: FAULT RECOVERY
 ─────────────────────────────────────────────────────────────────────
 Cloud preemption / GPU crash → auto-detect
 Load latest training checkpoint      ← resume from last save
 Verify model state + optimizer state ← integrity check
 Continue training seamlessly

         ↓

 PHASE 5: VALIDATION + DEPLOYMENT
 ─────────────────────────────────────────────────────────────────────
 Disable gradient checkpointing       ← NOT needed at inference!
 Re-enable use_cache                  ← needed for fast inference!
 model.eval() + torch.no_grad()       ← inference mode
 Export to ONNX / SafeTensors        ← deployment format
 Deploy as API endpoint               ← FastAPI / vLLM / TGI

 ═══════════════════════════════════════════════════════════════════════

 KEY MEMORY NUMBERS TO REMEMBER:
 Gradient Checkpointing alone:     saves ~60-80% activation memory
 + fp16 mixed precision:           saves another ~50% of model weights
 + DeepSpeed ZeRO Stage 2:         saves ~75% of optimizer states
 Combined effect:                  can train 5-10x larger models!

14. ✅❌ Best Practices — DOs and DON'Ts

✅ DO: Always use use_reentrant=False
The old reentrant checkpointing API has known bugs in distributed training and can produce incorrect gradients in some model architectures. Always pass use_reentrant=False in checkpoint() and gradient_checkpointing_enable(). This is the modern, stable API — use it for all new code!
❌ DON'T: Forget to disable KV cache during training
This is the #1 mistake beginners make with Hugging Face models. Always set model.config.use_cache = False before training. The attention KV cache is for fast inference — during training it conflicts with gradient checkpointing and wastes memory. Re-enable it after training: model.config.use_cache = True.
✅ DO: Combine gradient checkpointing with fp16/bf16 mixed precision
These two techniques attack different parts of memory — checkpointing reduces activations, mixed precision reduces weights and gradients. Together, they're far more powerful than either alone. The standard stack is: gradient checkpointing + bf16 (bfloat16 is more numerically stable than fp16 for modern GPUs).
❌ DON'T: Enable gradient checkpointing during inference
Gradient checkpointing only makes sense during training — when you need gradients for backpropagation. During inference (model.eval()), there is no backward pass, so activations don't need to be kept at all. PyTorch's torch.no_grad() context automatically handles this. Using checkpointing during inference just adds unnecessary computation overhead!
✅ DO: Save training checkpoints frequently with fault-tolerant design
In cloud environments (AWS Spot, GCP Preemptible, OCI Spot Instances), your GPU instance can be reclaimed anytime. Save a training checkpoint every 100–500 steps. Keep the last 3–5 checkpoints (circular buffer), not just the latest one. Always save the optimizer state AND the model state — without the optimizer state, your Adam optimizer restarts from scratch and training diverges!
❌ DON'T: Over-checkpoint (too many checkpoint segments)
More checkpoint segments doesn't always mean more memory savings. Research shows that exceeding the optimal number of checkpoints can actually increase memory consumption due to recomputation overhead and memory fragmentation. For Transformer models, the standard is: checkpoint each Transformer block (the combination of Attention + MLP). Don't go finer than that!
✅ DO: Profile memory before and after enabling checkpointing
Use torch.cuda.max_memory_allocated() before and after enabling gradient checkpointing. Confirm that memory usage actually dropped by the expected amount (typically 40–70%). If the savings are smaller than expected, check that use_cache=False is set and that the checkpointing was applied to the right layers!
❌ DON'T: Use gradient checkpointing as a substitute for proper model design
Gradient checkpointing trades computation for memory — it doesn't make your model more efficient overall. If you find yourself needing extreme checkpointing just to run a model, consider also: LoRA (reduce trainable parameters by 100x), quantization (QLoRA — 4-bit weights), or a smaller model architecture. These approaches are complementary, not alternatives!

15. 📝 Summary — Everything You Learned

The Core Concepts

  • Activations → Intermediate results computed at each layer during the forward pass. The biggest consumer of GPU memory during training.
  • Gradient Checkpointing → Discard activations during the forward pass. Recompute them during the backward pass. Saves ~60-80% activation memory at ~20-30% extra compute cost.
  • Memory complexity → Reduced from O(n) linear to O(√n) square root of layers. Massively reduces memory for deep models.
  • Training Checkpoint → A completely separate concept: saving the entire model+optimizer state to disk for fault tolerance and resumability.
  • use_reentrant=False → Always use the modern non-reentrant API. Safer, more compatible with distributed training.
  • use_cache=False → Always disable the KV attention cache during training when using gradient checkpointing.

Implementation Quick Reference

  • Custom PyTorch model: checkpoint(module_function, *inputs, use_reentrant=False)
  • Sequential model: checkpoint_sequential(model, segments=4, input=x, use_reentrant=False)
  • Hugging Face model: model.gradient_checkpointing_enable({"use_reentrant": False})
  • Hugging Face Trainer: gradient_checkpointing=True in TrainingArguments
  • DeepSpeed: "activation_checkpointing": {"partition_activations": true} in config JSON
  • FSDP + Selective: apply_activation_checkpointing(model, check_fn=lambda m: isinstance(m, TransformerBlock))

The Memory Optimization Stack

  • Step 1: Gradient Checkpointing → saves activation memory
  • Step 2: Mixed Precision (bf16) → saves weight and gradient memory
  • Step 3: Gradient Accumulation → allows larger effective batch sizes
  • Step 4: LoRA/QLoRA → reduces trainable parameters 100x
  • Step 5: FSDP / DeepSpeed ZeRO → shard across multiple GPUs

Layer by layer, this stack allows you to fine-tune 70B parameter models that previously required $100,000 GPU clusters — on a single A100!

Happy training! 🧠💾🚀

Comments