How to Implement Gradient Checkpointing for Memory-Efficient Large Model Training
Gradient checkpointing reduces GPU memory consumption by discarding intermediate activations during the forward pass and recomputing them during backpropagation, allowing you to train significantly larger transformer models on limited hardware.
Training large language models from scratch requires substantial GPU memory to store activations from every transformer layer. In the train-llm-from-scratch repository, implementing gradient checkpointing—also known as activation checkpointing—lets you trade modest computation overhead for dramatic memory savings, enabling you to scale model depth and batch size without upgrading your hardware.
What Is Gradient Checkpointing?
Gradient checkpointing is a memory optimization technique where selected forward activations are recomputed during the backward pass instead of being stored in GPU memory. In the transformer architecture implemented in src/models/transformer.py, the model consists of a stack of Block modules defined in src/models/transformer_block.py. Without checkpointing, every intermediate activation from attention scores to MLP hidden states remains in memory until backpropagation consumes them. By wrapping the forward pass of these blocks with torch.utils.checkpoint.checkpoint, you retain only the input tensors to each checkpointed region and recalculate the internals on demand.
This approach yields three primary effects: memory savings scale with the number of checkpointed layers, extra compute occurs because forward operations run twice (once during the initial pass and again during gradient computation), and transparency means the rest of your training loop in scripts/train_transformer.py requires no modifications.
Where to Apply Checkpointing in the Transformer Architecture
The most effective insertion point is inside the Transformer.forward method where the model iterates over self.attn_blocks. According to the source code, this loop sequentially processes each transformer block. By intercepting the plain call x = block(x) and replacing it with a checkpointed equivalent, you ensure that only the tensor x entering each block persists in memory, while all internal activations are temporary.
Implementation Methods
You can implement gradient checkpointing using three distinct approaches depending on whether you prefer fine-grained control, grouped efficiency, or non-invasive training script modifications.
Method 1: Checkpoint Individual Transformer Blocks
For maximum memory reduction, wrap each block individually. This requires modifying src/models/transformer.py to import torch.utils.checkpoint and adjust the forward loop:
# src/models/transformer.py
import torch.utils.checkpoint as checkpoint
class Transformer(nn.Module):
def forward(self, idx: torch.Tensor, targets: torch.Tensor = None):
x = self._pre_attn_pass(idx)
# Replace plain block calls with checkpointed calls
for block in self.attn_blocks:
x = checkpoint.checkpoint(block, x)
x = self.layer_norm(x)
logits = self.lm_head(x)
# ... loss calculation remains unchanged
return logits, loss
Effect: Memory usage drops roughly in proportion to the number of transformer blocks, as only the input x to each TransformerBlock is retained.
Method 2: Grouped Checkpointing with checkpoint_sequential
If you prefer to reduce recomputation overhead, group multiple blocks together using torch.utils.checkpoint.checkpoint_sequential. This checkpoints chunks of layers rather than individual units:
# src/models/transformer.py
from torch.utils.checkpoint import checkpoint_sequential
class Transformer(nn.Module):
def forward(self, idx: torch.Tensor, targets: torch.Tensor = None):
x = self._pre_attn_pass(idx)
# Configure chunk size based on your GPU memory/compute budget
chunksize = 2
block_chunks = [
self.attn_blocks[i:i + chunksize]
for i in range(0, len(self.attn_blocks), chunksize)
]
for chunk in block_chunks:
x = checkpoint_sequential(
nn.ModuleList(chunk),
len(chunk),
x
)
x = self.layer_norm(x)
logits = self.lm_head(x)
# ... loss calculation
return logits, loss
Effect: Fewer recomputations occur (one per chunk instead of one per block) while still achieving memory savings proportional to the number of chunks rather than layers.
Method 3: Training Script Wrapper (Non-Invasive)
To avoid modifying model source files, define a checkpointed forward function in scripts/train_transformer.py and replace the model's forward method at runtime:
# scripts/train_transformer.py
import torch.utils.checkpoint as checkpoint
def checkpointed_forward(idx, targets=None):
x = model._pre_attn_pass(idx)
# Apply checkpointing to each attention block
for block in model.attn_blocks:
x = checkpoint.checkpoint(block, x)
x = model.layer_norm(x)
logits = model.lm_head(x)
loss = None
if targets is not None:
B, T, C = logits.shape
loss = F.cross_entropy(
logits.view(B * T, C),
targets.view(B * T).long()
)
return logits, loss
# Replace model forward without altering src/models/
model.forward = checkpointed_forward
Effect: All checkpointing logic lives in your training script, keeping the core model implementation clean while achieving identical memory benefits.
Critical Implementation Details
When implementing gradient checkpointing in the train-llm-from-scratch codebase, consider these technical constraints to ensure stable training:
-
Determinism: Ensure checkpointed blocks contain pure functions without in-place tensor modifications or random state changes. The existing
TransformerBlockimplementation uses only deterministic operations, making it safe for checkpointing. -
Checkpoint Granularity: Fine-grained checkpointing (Method 1) maximizes memory savings but increases compute overhead. Coarse-grained grouping (Method 2) balances memory and speed. Choose granularity based on your GPU memory budget.
-
Device Placement: Checkpointing works transparently on both CPU and GPU. The repository already handles device placement via
config['device']in the training script, requiring no additional device management for checkpointed tensors. -
Mixed-Precision Compatibility: If using
torch.cuda.ampfor FP16/BF16 training, wrap checkpoint calls inside theautocastcontext manager. The checkpointing mechanism preserves gradient scaling behavior. -
torch.compile Limitations: PyTorch 2.0's
torch.compilemay not support gradient checkpointing in all configurations. Maintain eager mode for the model while using checkpoints to avoid graph compilation errors.
Summary
- Implement gradient checkpointing by wrapping
TransformerBlockforward calls insrc/models/transformer.pywithtorch.utils.checkpoint.checkpointto reduce memory usage proportional to your model depth. - Choose your granularity: Individual blocks maximize memory savings, while
checkpoint_sequentialgroups reduce recomputation overhead at the cost of higher peak memory. - Maintain purity: Ensure checkpointed regions contain no in-place operations or random state modifications to guarantee correct gradient flow.
- Verify compatibility: Use eager mode when combining checkpointing with
torch.compile, and wrap checkpoint calls inautocastcontexts for mixed-precision training.
Frequently Asked Questions
Does gradient checkpointing slow down training significantly?
Gradient checkpointing increases training time by approximately 10-30% depending on model architecture and checkpoint granularity. This overhead stems from recalculating forward passes during backpropagation. However, this trade-off enables training models that would otherwise require additional GPUs or gradient accumulation steps, often resulting in faster wall-clock time to convergence compared to CPU offloading or micro-batch strategies.
Can I use gradient checkpointing with torch.compile in PyTorch 2.0?
According to the source implementation, you should exercise caution when combining gradient checkpointing with torch.compile. While PyTorch 2.x improves compilation support, checkpointing may still cause graph breaks or unsupported operation errors in certain configurations. Until full compatibility is verified, run your model in eager mode when implementing checkpointing in the train-llm-from-scratch repository.
How much GPU memory can gradient checkpointing actually save?
Memory savings scale linearly with the number of checkpointed transformer blocks. For a model with N blocks, individual block checkpointing reduces activation memory from O(N) to O(1) (plus the memory for one block's activations during recomputation). In practice, this often allows training models 2-3x larger on the same GPU, or increasing batch size proportionally, making it essential for large-scale training runs.
Does checkpointing work with DistributedDataParallel (DDP)?
Yes, gradient checkpointing integrates seamlessly with PyTorch's DistributedDataParallel. The recomputed activations occur independently on each rank during the backward pass, and gradient synchronization happens normally after all-reduce operations. Ensure you apply checkpointing consistently across all ranks to maintain identical computational graphs and avoid distributed training deadlocks.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →