How to Implement Checkpointing with recompute_layer in Nanotron: A Complete Guide
Nanotron enables activation checkpointing via the recompute_layer configuration flag, which wraps transformer layer forward passes in PyTorch's checkpoint function to trade computation for reduced GPU memory usage.
Activation checkpointing is essential for training large language models with limited GPU memory. In Hugging Face's Nanotron framework, this feature is controlled through a simple configuration toggle that automatically applies gradient checkpointing to transformer layers without requiring changes to model architecture code.
Understanding the recompute_layer Mechanism
Nanotron implements activation checkpointing through the recompute_layer boolean flag defined in ParallelismConfig. When enabled, the framework wraps designated forward passes inside torch.utils.checkpoint.checkpoint, discarding intermediate activations after the forward pass and recomputing them during backpropagation.
The mechanism relies on a decorator-based approach for clean integration:
- Configuration – The
recompute_layerflag resides insrc/nanotron/config/parallelism_config.pyand defaults toFalse - Decoration – The
checkpoint_method(attr_name)decorator insrc/nanotron/utils.pyconditionally applies PyTorch checkpointing based on the instance attribute - Model Integration – Built-in models like LLaMA and Qwen read this flag during initialization and store it as
self.recompute_layer
Enabling Checkpointing in Your Configuration
To activate gradient checkpointing for your training run, modify your parallelism configuration:
from nanatron.config import ParallelismConfig
config = ParallelismConfig(
tp=2, # tensor parallelism
dp=4, # data parallelism
pp=1, # pipeline parallelism
recompute_layer=True, # Enable activation checkpointing
tp_recompute_allgather=True # Optional: recompute all-gather for tensor parallelism
)
Or in YAML configuration:
parallelism:
tp: 2
dp: 4
pp: 1
recompute_layer: true
tp_recompute_allgather: true
The tp_recompute_allgather flag (default True) provides additional memory savings for tensor-parallel training by recomputing the input all-gather operation during the backward pass, implemented in src/nanotron/parallel/tensor_parallel/functional.py.
Internal Implementation Details
The Configuration Layer
In src/nanotron/config/parallelism_config.py, the ParallelismConfig dataclass defines both checkpointing-related flags:
@dataclass
class ParallelismConfig:
recompute_layer: bool = False
tp_recompute_allgather: bool = True
# ... other parallelism settings
The Checkpoint Decorator
The core logic resides in src/nanotron/utils.py within the checkpoint_method decorator. This utility inspects the specified boolean attribute at runtime and, when True, forwards the method call to torch.utils.checkpoint.checkpoint:
from nanotron.utils import checkpoint_method
class TransformerLayer(nn.Module):
def __init__(self, config):
super().__init__()
self.recompute_layer = config.recompute_layer
self.attention = ...
self.mlp = ...
@checkpoint_method("recompute_layer")
def forward(self, hidden_states, attention_mask):
# Implementation automatically checkpointed when flag is True
return self.mlp(self.attention(hidden_states, attention_mask))
Model-Specific Integration
Built-in models implement this pattern consistently. In src/nanotron/models/llama.py (lines 740-774), the model stores the configuration flag and applies the decorator to the forward pass. Similarly, src/nanotron/models/qwen.py (lines 458-562) applies the same pattern, including handling for Mixture-of-Experts (MoE) layers.
The decorator expects the method's first argument to be self, with additional tensors passed positionally. It does not support *args or **kwargs.
Implementing Checkpointing in Custom Layers
To add activation checkpointing to custom modules in Nanotron, follow the established pattern from tests/test_checkpointing.py:
import torch
from torch import nn
from nanotron.utils import checkpoint_method
class CustomTransformerBlock(nn.Module):
def __init__(self, dim, recompute: bool = False):
super().__init__()
self.norm = nn.LayerNorm(dim)
self.linear = nn.Linear(dim, dim)
self.recompute_layer = recompute # Required attribute name
@checkpoint_method("recompute_layer")
def forward(self, x: torch.Tensor, position_ids: torch.Tensor):
# Must use positional arguments only
normalized = self.norm(x)
return self.linear(normalized)
Important constraints:
- The decorated method must use positional arguments only (no
*argsor**kwargs) - The attribute name passed to the decorator must match the instance variable exactly
- The first parameter must be
self
Interactions with Parallelism Features
Tensor Parallelism – When tp_recompute_allgather is enabled alongside recompute_layer, Nanotron recomputes the tensor-parallel all-gather operations during backpropagation. This reduces activation memory at the cost of additional communication computation during the backward pass.
Pipeline Parallelism – Checkpointing operates transparently across pipeline stages. Each stage recomputes its own forward pass independently during backpropagation, maintaining the same pipeline bubble characteristics while reducing per-stage memory consumption.
Memory vs. Speed Trade-off – Enabling recompute_layer typically increases training step time by 20-30% (depending on model size and hardware) while reducing activation memory by approximately 50%, enabling training of models that would otherwise exceed GPU memory limits.
Summary
- Configuration-driven: Set
recompute_layer=TrueinParallelismConfigto enable checkpointing without code changes - Decorator-based: The
checkpoint_methoddecorator insrc/nanotron/utils.pywraps forward passes conditionally based on the configuration flag - Model support: Built-in implementations in
src/nanotron/models/llama.pyandsrc/nanotron/models/qwen.pydemonstrate proper integration patterns - Tensor-parallel optimization: Use
tp_recompute_allgatherfor additional memory savings in tensor-parallel setups - Custom modules: Apply
@checkpoint_method("recompute_layer")to custom layers, ensuring positional arguments only - Testing: Verify behavior using
tests/test_checkpointing.pywhich confirms forward passes execute twice (once for forward, once for recompute) while storing no intermediate activations
Frequently Asked Questions
What is activation checkpointing and why should I use it?
Activation checkpointing is a memory optimization technique that trades computation for memory by discarding intermediate activations during the forward pass and recomputing them during backpropagation. Use it when training models larger than approximately 30 billion parameters or when GPU memory constraints prevent increasing batch size or sequence length.
Does enabling recompute_layer slow down training?
Yes, enabling recompute_layer increases step time because each checkpointed layer's forward pass executes twice—once during the initial forward pass and again during backpropagation. However, this trade-off is often necessary to fit larger models into memory or use larger batch sizes, which can improve overall training throughput and model convergence.
Can I use recompute_layer with custom models not in the Nanotron repository?
Absolutely. Any nn.Module can support checkpointing by adding a boolean attribute (conventionally named recompute_layer) and decorating the forward method with @checkpoint_method("recompute_layer") from src/nanotron/utils.py. Ensure your forward method accepts only positional arguments besides self.
What is the difference between recompute_layer and tp_recompute_allgather?
recompute_layer controls checkpointing of transformer layer computations (attention and MLP blocks), while tp_recompute_allgather specifically controls recomputation of the tensor-parallel input all-gather operations. The former reduces memory from layer activations; the latter reduces memory from tensor-parallel communication buffers. Both default to complementary usage but can be configured independently based on your memory and performance requirements.
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 →