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 TransformerBlock implementation 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.amp for FP16/BF16 training, wrap checkpoint calls inside the autocast context manager. The checkpointing mechanism preserves gradient scaling behavior.

  • torch.compile Limitations: PyTorch 2.0's torch.compile may 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 TransformerBlock forward calls in src/models/transformer.py with torch.utils.checkpoint.checkpoint to reduce memory usage proportional to your model depth.
  • Choose your granularity: Individual blocks maximize memory savings, while checkpoint_sequential groups 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 in autocast contexts 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:

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 →