# How to Implement Gradient Checkpointing for Memory-Efficient Large Model Training

> Implement gradient checkpointing to reduce GPU memory for large model training. Train bigger transformer models with less hardware by recomputing activations during backpropagation.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: how-to-guide
- Published: 2026-05-31

---

**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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py), the model consists of a stack of `Block` modules defined in [`src/models/transformer_block.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py) to import `torch.utils.checkpoint` and adjust the forward loop:

```python

# 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:

```python

# 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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py) and replace the model's `forward` method at runtime:

```python

# 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`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/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.