How Activation Checkpointing in MegaDLMs Reduces Memory Usage

Activation checkpointing in MegaDLMs trades additional forward computation for reduced GPU memory by discarding intermediate transformer activations during the forward pass and recomputing them on-the-fly during backward propagation.

MegaDLMs (Megatron-based Deep Language Models) implement activation checkpointing to train massive transformer models that would otherwise exceed GPU memory limits. This technique, configured through megatron/core/transformer/transformer_config.py and executed in megatron/core/transformer/transformer_block.py, can reduce per-GPU activation memory by up to 70% while increasing training time by only 10-15%.

What Is Activation Checkpointing in MegaDLMs?

During standard transformer training, every intermediate activation—such as query-key-value projections, attention scores, and feed-forward network outputs—is retained in GPU memory until the backward pass completes. For deep models with billions of parameters, these activations dominate the memory footprint.

Activation checkpointing in MegaDLMs modifies this behavior by retaining only a minimal set of checkpoint tensors (typically the inputs to transformer layers or attention blocks) and discarding all other intermediate values. When the backward pass requires a discarded activation, the forward computation for that specific segment is re-executed to regenerate the needed values. This recomputation costs additional FLOPs but eliminates the need to store the full activation graph.

Implementation Details in the MegaDLMs Codebase

Configuration Parameters in transformer_config.py

The checkpointing behavior is controlled through the TransformerConfig class in megatron/core/transformer/transformer_config.py. Key fields include:

  • recompute_granularity: Set to "selective" to checkpoint only the memory-intensive attention components, or "full" for complete layer checkpointing.
  • recompute_method: Specifies "uniform" to split the model into equal chunks, or "block" to checkpoint a specific number of consecutive layers.
  • recompute_num_layers: Defines the number of layers per checkpoint group when using uniform or block methods.

As documented in lines 182-186 of the configuration file:


# megatron/core/transformer/transformer_config.py

recompute_granularity: str = None
"""'selective' activation checkpointing where only the memory intensive part of
attention is checkpointed. 'full' checkpoints the entire transformer layer."""

The Checkpointed Forward Pass in transformer_block.py

The core logic resides in megatron/core/transformer/transformer_block.py within the _checkpointed_forward method (starting at line 275). This function constructs a custom forward handler that integrates with either the Tensor-Parallel checkpoint wrapper or Transformer-Engine's checkpointing depending on the FP8 configuration:


# megatron/core/transformer/transformer_block.py

def _checkpointed_forward(self, hidden_states, attention_mask, ...):
    """Forward method with activation checkpointing."""
    
    def custom_forward(*inputs):
        # Unpack inputs and run forward

        ...
        return output
    
    def checkpoint_handler(forward_func):
        if self.config.fp8:
            return te_checkpoint(forward_func, ...)
        else:
            return tensor_parallel.checkpoint(forward_func, ...)
    
    # Apply checkpointing

    return checkpoint_handler(custom_forward)(hidden_states, ...)

Selective vs. Full Checkpointing Strategies

MegaDLMs supports two granularity levels for activation checkpointing:

Selective checkpointing targets only the attention mechanism's query-key-value projections and softmax computations—the largest activation tensors in transformer layers—while retaining feed-forward network activations. This approach minimizes recomputation overhead while capturing the majority of memory savings.

Full checkpointing stores only the layer inputs and discards all internal activations, requiring complete layer recomputation. This maximizes memory reduction but increases training time by approximately 20-30% compared to selective checkpointing.

Uniform vs. Block Recompute Methods

The distribution of checkpointed layers is controlled by the recompute_method parameter in transformer_block.py (lines 332-465):

  • Uniform method: Divides the total number of layers into equal-sized chunks defined by recompute_num_layers. The input to each chunk is checkpointed, allowing efficient parallel recomputation across chunks during the backward pass.

  • Block method: Applies checkpointing to a specific contiguous block of recompute_num_layers layers, while remaining layers run without checkpointing. This is useful when only a subset of layers (typically early or late layers) exhibit high memory pressure.

Memory Savings and Performance Impact

Activation checkpointing fundamentally changes the memory-compute trade-off in distributed training:

Training Phase Memory Without Checkpointing Memory With Checkpointing Compute Overhead
Forward Pass Stores all intermediate activations (Q, K, V, FFN outputs) Stores only checkpoint inputs None
Backward Pass Uses cached activations Recomputes forward for each checkpointed segment ~10-15% additional FLOPs

For a 30-billion parameter model running on NVIDIA A100 GPUs, enabling selective activation checkpointing reduces per-GPU activation memory from approximately 30 GB to 10 GB—a 70% reduction—while increasing total training time by only 10-15%. This makes it feasible to train larger models on single 40 GB A100 GPUs without requiring excessive model parallelism across dozens of devices.

Configuring Activation Checkpointing in Training Scripts

To enable activation checkpointing in MegaDLMs, modify the TransformerConfig before model initialization:

from megatron.core.transformer.transformer_config import TransformerConfig

# Configure checkpointing parameters

config = TransformerConfig(
    num_layers=32,
    hidden_size=4096,
    num_attention_heads=32,
    recompute_granularity="selective",      # "selective" or "full"

    recompute_method="uniform",               # "uniform" or "block"

    recompute_num_layers=2,                   # Layers per checkpoint group

    distribute_saved_activations=True,        # Distribute across model parallel group

    fp8=None,                                 # Set to "e4m3" to use Transformer-Engine

)

When using the higher-level training API, these parameters are typically passed through command-line arguments or Hydra configuration files. The TransformerBlock automatically detects these settings and routes the forward pass through _checkpointed_forward when checkpointing is enabled.

To verify memory savings during training:

import torch
from megatron.training import training

# Build model with checkpointing enabled

model = training.build_model(config)

# Monitor memory

torch.cuda.reset_peak_memory_stats()
output = model(input_batch)
loss = output.mean()
loss.backward()
print(f"Peak memory: {torch.cuda.max_memory_allocated() / 1e9:.2f} GB")

Summary

  • Activation checkpointing in MegaDLMs reduces GPU memory by discarding intermediate transformer activations during the forward pass and recomputing them on-the-fly during backward propagation.
  • The implementation in megatron/core/transformer/transformer_block.py supports selective (attention-only) and full layer checkpointing, controlled via TransformerConfig parameters defined in transformer_config.py.
  • Uniform and block recompute methods provide flexible strategies for distributing checkpointing across layers to optimize memory fragmentation and computational efficiency.
  • Enabling selective checkpointing typically reduces activation memory by ~70% (for example, from 30 GB to 10 GB per GPU) with only 10-15% additional compute overhead, enabling larger models to fit on existing hardware.

Frequently Asked Questions

What is the difference between selective and full activation checkpointing in MegaDLMs?

Selective checkpointing targets only the memory-intensive components of the attention mechanism—specifically the query-key-value projections and softmax operations—while retaining feed-forward network activations. Full checkpointing stores only the layer inputs and discards all internal activations, requiring complete layer recomputation. Selective checkpointing offers the best memory-to-compute trade-off for most transformer workloads, while full checkpointing is used when memory constraints are most severe.

How much memory does activation checkpointing actually save when training large models?

According to the MegaDLMs implementation, enabling selective activation checkpointing reduces per-GPU activation memory by approximately 70%. For example, a 30-billion parameter model that normally requires 30 GB of activation memory per GPU consumes only 10 GB when checkpointing is enabled, allowing training on single 40 GB A100 GPUs without requiring additional model parallelism across dozens of devices.

Does activation checkpointing slow down training significantly?

Activation checkpointing trades compute for memory, but the overhead is modest. The MegaDLMs codebase indicates that selective checkpointing increases total training time by approximately 10-15%, while full checkpointing may add 20-30% overhead. This is typically acceptable given the alternative of running out of GPU memory or requiring additional expensive hardware, and the overhead decreases relative to total training time as model size increases.

Where is the activation checkpointing logic implemented in the MegaDLMs repository?

The core logic resides in megatron/core/transformer/transformer_block.py within the _checkpointed_forward method (lines 275-465), which orchestrates the checkpointed segments. Configuration parameters are defined in megatron/core/transformer/transformer_config.py (lines 182-186). The actual checkpointing wrapper is provided by tensor_parallel.checkpoint in megatron/core/tensor_parallel/random.py when FP8 is disabled, or te_checkpoint from Transformer-Engine when FP8 is enabled.

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 →