Trading Compute for Memory: Using torch.utils.checkpoint in PyTorch
Stop hitting CUDA out-of-memory errors. Learn how to use torch.utils.checkpoint to trade a bit of compute time for significant GPU memory savings in deep PyTorch models.
17 Dec 2025, 02:50 UTC

The Out-of-Memory Wall
When training deep architectures—such as large Transformers or deep ResNets—you often hit a GPU memory ceiling. This isn't usually caused by the model weights themselves, but by activations. PyTorch stores every intermediate tensor produced during the forward pass because they are required to calculate gradients during the backward pass. As your model depth or batch size increases, these activations can easily exceed your available VRAM, leading to the dreaded RuntimeError: CUDA out of memory.
The Thesis: Activation Checkpointing
The most effective way to break through this ceiling without upgrading hardware is activation checkpointing via torch.utils.checkpoint. The strategy is simple: instead of storing every intermediate activation, PyTorch discards them during the forward pass and re-computes them on-the-fly during the backward pass. You trade a small amount of additional computation time for a significant reduction in peak memory usage.
Implementing Checkpointing in a Model
Checkpointing is applied by wrapping a segment of the model's forward pass. It is most effective when applied to repeated blocks (like ResNet bottlenecks or Transformer layers) rather than the entire model.
import torch
import torch.nn as nn
import torchvision.models as models
from torch.utils.checkpoint import checkpoint
# A wrapper to integrate checkpointing into a standard nn.Module
class CheckpointedBlock(nn.Module):
def __init__(self, block):
super().__init__()
self.block = block
def forward(self, x):
# checkpoint() takes a callable and the arguments for that callable
# It prevents the storage of intermediate activations within self.block
return checkpoint(self.block, x, use_reentrant=False)
# Setup: Using ResNet-50 as a target for memory reduction
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = models.resnet50(weights=None).to(device)
# Selectively wrap the second layer's blocks to save memory
for i in range(len(model.layer2)):
model.layer2[i] = CheckpointedBlock(model.layer2[i])
# Verification Setup
inputs = torch.randn(32, 3, 224, 224, device=device)
targets = torch.randint(0, 1000, (32,), device=device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
# Execution
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step()
Technical Constraints and Risks
- Compute Overhead: Because the forward pass for checkpointed sections is executed twice, you will see an increase in training time per iteration. Depending on the model, this typically ranges from 10% to 30%.
- Randomness (Dropout/BatchNorm): If a checkpointed block contains stochastic operations like Dropout, PyTorch must ensure the random seed is preserved so the re-computed activation matches the original.
torch.utils.checkpointhandles this internally, but custom random logic may require manual seeding. - In-place Modifications: Avoid using in-place operations (e.g.,
relu_(x)) inside a checkpointed function. Since the original input is needed for re-computation, modifying it in-place can lead to incorrect gradients or runtime errors.
Measuring the Impact
- Memory Baseline: Use
torch.cuda.memory_summary()before and after applying checkpointing to observe the drop in peak allocated memory. - Timing Overhead: Wrap your training loop in a timer. If your forward pass is not the primary bottleneck, the 15-20% time increase is often a fair trade for doubling your batch size.
- Numerical Parity: Use a fixed seed and compare the loss values of a non-checkpointed run versus a checkpointed run. They should be identical within floating-point precision.
Rollback Procedure
To revert to standard training, remove the CheckpointedBlock wrappers and replace them with the original nn.Module instances. This will restore maximum training speed but will return the model to its original peak memory consumption.
0 replies
A thoughtful contribution can make all the difference. Be the first to share one.