Memory Management: Activations, Gradients

Michael BrenndoerferJanuary 16, 202647 min read

Part of Language AI Handbook

Explains how GPU memory breaks down into parameters, gradients, optimizer states, and activations. Estimate memory requirements and debug out-of-memory errors.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

Memory Management

Every large language model training run is ultimately constrained by one resource before compute, before data, even before time: GPU memory. A model that fits in memory trains; one that does not, crashes. Understanding exactly what occupies GPU memory, how much each component costs, and how to diagnose failures when things go wrong separates practitioners who can scale models from those who hit walls.

This chapter breaks down the four major memory consumers in transformer training: model parameters, gradients, optimizer states, and activations. For each category, you will learn how much memory it requires, why that amount is unavoidable, and what tradeoffs appear when you try to reduce it. You will then build a working memory estimator that predicts whether a model fits in a given GPU budget before you ever start training, and finally work through systematic OOM (out-of-memory) debugging strategies to recover from crashes when estimation alone is not enough.

The concepts here connect directly to the distributed training strategies covered in the next chapters. Tensor parallelism, pipeline parallelism, and gradient checkpointing all exist primarily to manage memory. Without understanding what they are managing, you cannot reason about when to use them or what tradeoffs they impose.

The Four Memory Consumers

GPU memory during training is divided among four distinct categories, each with different characteristics and different strategies for reduction. The total memory budget must accommodate all four simultaneously, and failing to account for any one of them is a common source of mysterious OOM errors late in training runs.

One useful framing: think of parameter memory, gradient memory, and optimizer state memory as static costs that depend only on model architecture. They exist as soon as the model and optimizer are instantiated and they persist until training ends. Activation memory, by contrast, is a dynamic cost that depends on the current batch: its size, the sequence length, and whether you have enabled techniques to reduce it. The two categories require different strategies for reduction and different tools for debugging.

Parameters

Model parameters are the learned weights stored in the model: embedding matrices, attention projection matrices, feed-forward network weights, layer normalization scales and biases, and output projection weights. These are the values that exist before training begins (initialized randomly or from a checkpoint) and persist throughout the entire training process.

For a transformer with PP parameters, the parameter memory in bytes depends directly on the numeric precision used:

  • In full 32-bit float (FP32): 4P4P bytes
  • In 16-bit half-precision float (FP16 or BF16): 2P2P bytes
  • In 8-bit integer (INT8): PP bytes

Modern large-scale training almost universally uses BF16 or FP16 for parameters during the forward pass, so a model with 7 billion parameters requires roughly 2×7×109=142 \times 7 \times 10^9 = 14 GB just for parameters. A 70B model requires 140 GB. These numbers set a hard floor: you need at least this much memory before accounting for anything else.

The count of parameters PP in a transformer scales predictably. For a model with:

  • Vocabulary size VV
  • Hidden dimension dd
  • Number of layers LL
  • Feedforward intermediate dimension dffd_{ff} (typically 4d4d)
  • Number of attention heads hh, with head dimension dh=d/hd_h = d/h

The dominant parameter contributions are:

P≈V⋅d+L⋅(4d2+2⋅d⋅dff)+dP \approx V \cdot d + L \cdot \left(4d^2 + 2 \cdot d \cdot d_{ff}\right) + d

where the first term covers the embedding table, the second covers attention projections and feed-forward weights per layer, and the third covers the final layer norm. The attention projection contribution of 4d24d^2 comes from four projection matrices WQW_Q, WKW_K, WVW_V, WOW_O, each of shape d×dd \times d. The feedforward contribution 2⋅d⋅dff2 \cdot d \cdot d_{ff} comes from two linear layers: one expanding from dd to dffd_{ff} and one projecting back to dd.

To make this concrete, for a 7B-parameter model with d=4096d = 4096, L=32L = 32, dff=11,008d_{ff} = 11{,}008, V=32,000V = 32{,}000:

Embeddings=32,000×4096≈131MPer-layer attention=4×40962≈67MPer-layer FFN=2×4096×11,008≈90MPer-layer total≈157MAll 32 layers≈5,024MGrand total≈6,700M≈7B\begin{aligned} \text{Embeddings} &= 32{,}000 \times 4096 \approx 131M \\ \text{Per-layer attention} &= 4 \times 4096^2 \approx 67M \\ \text{Per-layer FFN} &= 2 \times 4096 \times 11{,}008 \approx 90M \\ \text{Per-layer total} &\approx 157M \\ \text{All 32 layers} &\approx 5{,}024M \\ \text{Grand total} &\approx 6{,}700M \approx 7B \end{aligned}

This derivation shows how the bulk of parameters live in the transformer layers themselves. The embedding table is a significant contributor for models with large vocabularies, but for modern decoder-only models the layer weights dominate.

BF16 vs FP16

Both BF16 and FP16 use 16 bits, but they partition those bits differently. FP16 allocates 10 bits to the mantissa (fractional precision) and 5 bits to the exponent (range). BF16 allocates only 7 bits to the mantissa but 8 bits to the exponent, matching the range of FP32. This matters because training involves large gradient magnitudes that can overflow FP16's limited range, causing numerical instability. BF16 avoids this overflow problem at the cost of slightly lower precision, making it the preferred choice for training large transformers. When you see a model described as "BF16 training," the parameters, forward pass activations, and gradient accumulation all happen in BF16 except where FP32 is explicitly required for numerical stability.

Gradients

During backpropagation, the training loop computes the gradient of the loss with respect to every parameter. These gradients have exactly the same shape as the parameters themselves, meaning they require an identical amount of memory. If you store parameters in FP16 (2P2P bytes), you need additional memory for the corresponding gradients.

The precision question for gradients is subtler than for parameters. In mixed-precision training, the forward pass runs in FP16 or BF16 to take advantage of fast half-precision matrix multiplication hardware. But accumulating small gradient signals across millions of samples in FP16 is numerically dangerous: the limited mantissa precision causes small gradients to underflow to zero and large ones to overflow. The standard solution is to store and accumulate gradients in FP32, which provides enough dynamic range and precision to keep the update signal intact. This means gradients often cost 4P4P bytes even when parameters themselves are stored in FP16.

In practice, the mixed-precision pipeline looks like this: during the forward pass, parameters are cast from their FP32 master copies to FP16 for the actual computation. The resulting FP16 activations flow through the network. During the backward pass, gradients are computed in FP16 and then immediately cast to FP32 before accumulation. The optimizer update then operates entirely in FP32. After the update, the updated FP32 parameters are cast back to FP16 for the next forward pass. The FP32 master copy exists to accumulate small updates that would otherwise vanish in FP16.

Gradients exist only during backpropagation. They do not accumulate across batches (unless you are using gradient accumulation, in which case you are intentionally summing them). After the optimizer applies the update and zero_grad() is called, this memory is freed. However, at peak training memory usage, the gradient tensors are fully populated and occupy their full allocation.

Gradient memory can be reduced through gradient compression (sending lower-precision gradients in distributed settings), gradient accumulation (which does not reduce peak memory but allows using smaller batches), and gradient checkpointing (which trades compute for activation memory, not directly for gradient memory). The gradient footprint itself is fundamentally tied to parameter count.

Optimizer States

The optimizer is often the largest single memory consumer, exceeding both parameters and gradients. This surprises many practitioners who focus on model size without accounting for the optimizer.

Adam, the most commonly used optimizer for transformer training, maintains two moving average tensors for every parameter:

  • The first moment estimate mtm_t (exponentially decaying average of gradients)
  • The second moment estimate vtv_t (exponentially decaying average of squared gradients)

The Adam update rule ties these together. Given the gradient gtg_t at step tt, Adam computes:

mt=β1mt−1+(1−β1)gtvt=β2vt−1+(1−β2)gt2m^t=mt1−β1t(bias correction)v^t=vt1−β2t(bias correction)θt=θt−1−ηv^t+ϵm^t\begin{aligned} m_t &= \beta_1 m_{t-1} + (1 - \beta_1) g_t \\ v_t &= \beta_2 v_{t-1} + (1 - \beta_2) g_t^2 \\ \hat{m}_t &= \frac{m_t}{1 - \beta_1^t} \quad \text{(bias correction)} \\ \hat{v}_t &= \frac{v_t}{1 - \beta_2^t} \quad \text{(bias correction)} \\ \theta_t &= \theta_{t-1} - \frac{\eta}{\sqrt{\hat{v}_t} + \epsilon} \hat{m}_t \end{aligned}

where:

  • β1≈0.9\beta_1 \approx 0.9 is the decay rate for the first moment
  • β2≈0.999\beta_2 \approx 0.999 is the decay rate for the second moment
  • η\eta is the learning rate
  • ϵ\epsilon is a small constant for numerical stability (typically 10−810^{-8})

Both mtm_t and vtv_t are typically stored in FP32, regardless of parameter precision. For a model with PP parameters, Adam requires:

Adam optimizer memory=4P+4P=8P bytes\text{Adam optimizer memory} = 4P + 4P = 8P \text{ bytes}

This means that for a 7B parameter model, Adam alone requires approximately 8×7×109=568 \times 7 \times 10^9 = 56 GB. Combined with FP16 parameters (1414 GB) and FP32 gradients (2828 GB), the total before activations is 14+28+56=9814 + 28 + 56 = 98 GB for a 7B model trained with Adam in mixed precision.

The reason both moments are stored in FP32 rather than FP16 is numerical stability. The second moment estimate vtv_t accumulates squared gradient values, and the Adam update divides by vt\sqrt{v_t}. Very small values of vtv_t in low-precision arithmetic lead to division instability and divergence. The first moment mtm_t is similarly vulnerable: if the gradient signal is small and mtm_t underflows to zero in FP16, the optimizer loses the ability to update parameters with slowly changing gradients. FP32 provides enough dynamic range to prevent both failure modes.

Alternative Optimizers

When Adam's 8P8P bytes is too expensive, several alternatives offer lower memory overhead with varying tradeoffs:

Adafactor approximates the second moment using a factored representation. For a weight matrix WW of shape m×nm \times n, Adam stores a full m×nm \times n second moment matrix. Adafactor instead stores a single row vector of size mm and a column vector of size nn, whose outer product approximates the full matrix. This reduces memory from O(mn)O(mn) to O(m+n)O(m + n) for each matrix. Across the model, optimizer memory drops from O(P)O(P) to roughly O(P)O(\sqrt{P}). Adafactor was used to train T5 and many subsequent encoder-decoder models and works well when the effective learning rate schedule is managed carefully.

Lion (Evolved Sign Momentum) uses only the sign of the gradient update, maintaining a single exponential moving average of gradients rather than two. The update rule is:

ct=sign(β1mt−1+(1−β1)gt)mt=β2mt−1+(1−β2)gtθt=θt−1−η⋅ct\begin{aligned} c_t &= \text{sign}(\beta_1 m_{t-1} + (1 - \beta_1) g_t) \\ m_t &= \beta_2 m_{t-1} + (1 - \beta_2) g_t \\ \theta_t &= \theta_{t-1} - \eta \cdot c_t \end{aligned}

Lion requires only one FP32 state tensor per parameter, halving optimizer memory from 8P8P to 4P4P bytes. Because it applies uniform update magnitude (scaled by learning rate), it is effectively a sign-gradient method and needs a lower learning rate than Adam to achieve comparable results.

SGD with Momentum stores only the velocity buffer (4P4P bytes in FP32), consuming half the memory of Adam. It is rarely used for transformer pretraining because it converges more slowly and is more sensitive to learning rate choices. However, it can be effective for fine-tuning tasks where the model is already near a good solution.

8-bit Adam quantizes optimizer states to 8-bit integers using dynamic quantization. Each block of optimizer state values is stored with a per-block scaling factor. This reduces the 8P8P bytes of standard Adam to 2P2P bytes while maintaining similar training dynamics. The bitsandbytes library provides a drop-in replacement that works with most PyTorch training code.

Mixed Precision Memory Formula

The standard mixed precision training setup (FP16 parameters, FP32 gradients, FP32 Adam states) requires:

Total non-activation memory=2P+4P+8P=14P bytes\text{Total non-activation memory} = 2P + 4P + 8P = 14P \text{ bytes}

This "14x rule" is widely used in practice: multiply the parameter count in billions by 14 to get the minimum GPU memory in gigabytes needed before activations. A 7B model requires about 98 GB, a 13B model requires about 182 GB, and a 70B model requires close to 1 TB. These numbers immediately explain why multi-GPU and multi-node training is mandatory for large language model development.

Activations

Activations are the intermediate tensors computed during the forward pass, stored in memory so that backpropagation can compute gradients. Unlike parameters, gradients, and optimizer states, activation memory scales with model size, batch size, and sequence length. This dynamic behavior makes activation memory the hardest component to predict and the one most often responsible for late-stage OOM crashes.

What Activations Are Stored

To understand why so much memory is needed for activations, it helps to walk through exactly what needs to be stored for a single transformer layer during the backward pass.

Consider a single transformer layer processing an input tensor XX of shape (B,T,d)(B, T, d), where BB is batch size, TT is sequence length, and dd is the hidden dimension. The forward pass computes:

  1. Layer norm pre-attention: The normalized output X^1\hat{X}_1 of shape (B,T,d)(B, T, d) must be retained to compute the layer norm gradient.
  2. Query, Key, Value projections: The tensors Q=X^1WQQ = \hat{X}_1 W_Q, K=X^1WKK = \hat{X}_1 W_K, V=X^1WVV = \hat{X}_1 W_V, each of shape (B,T,d)(B, T, d), must be retained to compute projection weight gradients.
  3. Attention score matrix: The raw scores S=QK⊤/dhS = QK^\top / \sqrt{d_h} of shape (B,h,T,T)(B, h, T, T) must be retained for the softmax gradient.
  4. Attention weights after softmax: The softmax output A=softmax(S)A = \text{softmax}(S) of shape (B,h,T,T)(B, h, T, T) must be retained for the value gradient.
  5. Attention output: The attended values AVAV of shape (B,T,d)(B, T, d) must be retained.
  6. Output projection input: The pre-projection activations for WOW_O.
  7. Layer norm pre-FFN: The normalized input to the feedforward block.
  8. FFN intermediate activations: The output of the first linear layer (before the activation function) of shape (B,T,dff)(B, T, d_{ff}).
  9. FFN activation output: The post-GELU or post-ReLU tensor, shape (B,T,dff)(B, T, d_{ff}).
  10. Layer norm values: The mean and variance used in layer normalization must be stored to compute layer norm gradients.

The attention weight matrices at items 3 and 4 are particularly expensive. Each of the LL layers stores two attention weight matrices of shape (B,h,T,T)(B, h, T, T). At FP16 (2 bytes), each matrix costs:

2×B×h×T2 bytes2 \times B \times h \times T^2 \text{ bytes}

For a 32-layer model with B=4B=4, h=32h=32, T=2048T=2048:

2×4×32×20482≈2.1 GB per attention matrix2 \times 4 \times 32 \times 2048^2 \approx 2.1 \text{ GB per attention matrix}

With two matrices (scores and weights) and 32 layers, the attention-only activation memory is approximately 134 GB. This is where the quadratic term dominates.

The Activation Memory Formula

Summing across all components of a single transformer layer and LL layers, the activation memory for standard full storage is approximately:

Mact≈L⋅B⋅T⋅(34d+5hT) bytesM_{\text{act}} \approx L \cdot B \cdot T \cdot \left(34d + 5 h T\right) \text{ bytes}

where the factor of 2 bytes per element (FP16) is absorbed into the constants. The first term 34d34d captures all the linear activations (projections, layer norms, FFN intermediate tensors), and the second term 5hT5hT captures the quadratic attention weight matrices. The T2T^2 dependence in the second term means:

  • Doubling sequence length quadruples the attention activation cost
  • For long sequences, attention activations dominate over all other activation memory

For a 7B parameter GPT-style model with d=4096d=4096, L=32L=32, h=32h=32, B=4B=4, and T=2048T=2048:

Mact≈32×4×2048×(34×4096+5×32×2048)=32×4×2048×(139,264+327,680)=32×4×2048×466,944≈122 GB\begin{aligned} M_{\text{act}} &\approx 32 \times 4 \times 2048 \times (34 \times 4096 + 5 \times 32 \times 2048) \\ &= 32 \times 4 \times 2048 \times (139{,}264 + 327{,}680) \\ &= 32 \times 4 \times 2048 \times 466{,}944 \\ &\approx 122 \text{ GB} \end{aligned}

This activation footprint nearly matches the static memory cost. At longer sequence lengths or larger batch sizes, activations quickly become the dominant consumer.

Gradient Checkpointing

Gradient checkpointing directly addresses the O(L×B×T)O(L \times B \times T) activation memory problem. The key insight is that activations are needed for backpropagation but not for the forward pass itself. If you are willing to recompute them during the backward pass, you can discard them after the forward pass and only keep enough state to perform the recomputation.

In its simplest form, gradient checkpointing works as follows:

  1. During the forward pass: Do not store activations for backpropagation. Instead, store only the input to each checkpointed segment (typically one transformer layer or a small group of layers).
  2. During the backward pass: When the backward algorithm needs the activations of a layer, rerun the forward pass for that segment from its stored input. Use the recomputed activations to compute gradients, then discard them again.

The memory trade is: instead of storing all LL layers' activations simultaneously, you store only one layer's activations at any given time (the one currently being processed in the backward pass). This reduces activation memory from O(L×B×T×d)O(L \times B \times T \times d) to approximately O(B×T×d)O(B \times T \times d).

The compute cost is additional forward passes. For each layer, you run the forward computation twice: once during the original forward pass, and once during the backward pass. Across all LL layers, this means roughly L+L=2LL + L = 2L forward passes instead of LL, an overhead of approximately 33% more compute (since the backward pass itself is roughly as expensive as the forward pass, the total goes from 1F+1B1F + 1B to 1F+L⋅1LF+1B=2F+1B1F + L \cdot \frac{1}{L}F + 1B = 2F + 1B, which for F≈BF \approx B means about 33% more total compute).

PyTorch implements this through torch.utils.checkpoint.checkpoint(). For HuggingFace models, the convenience method model.gradient_checkpointing_enable() applies checkpointing to every transformer layer. The tradeoff is always worth it when you are activation-memory-limited: 33% more compute to save 10-100x activation memory is a favorable exchange in most training scenarios.

A more granular strategy is selective checkpointing: store activations for computationally cheap operations (like layer norms) but recompute activations for expensive operations (like attention matrices). FlashAttention, which we cover in Part 14, takes this idea further by restructuring the attention computation to never fully materialize the T×TT \times T attention weight matrix, saving both memory and bandwidth.

Memory Estimation

Estimating total memory before running a training job lets you choose the right GPU configuration, set batch sizes appropriately, and avoid wasting hours waiting for an OOM crash that was predictable from the start. A good estimate also helps you make principled decisions about which memory reduction techniques to apply and in what order.

The Estimation Formula

For a standard mixed-precision training setup with Adam, the total memory MM breaks down as:

M=2P⏟FP16 params+4P⏟FP32 grads+8P⏟Adam states+Mact⏟activations+Mother⏟overheadM = \underbrace{2P}_{\text{FP16 params}} + \underbrace{4P}_{\text{FP32 grads}} + \underbrace{8P}_{\text{Adam states}} + \underbrace{M_{\text{act}}}_{\text{activations}} + \underbrace{M_{\text{other}}}_{\text{overhead}}

where MotherM_{\text{other}} covers CUDA context, framework overhead, and temporary buffers. The CUDA runtime itself requires roughly 300-500 MB when initialized. PyTorch maintains a memory caching allocator that retains freed memory for reuse; its overhead depends on the number and sizes of allocations. Temporary workspace for cuBLAS and cuDNN operations can require several hundred megabytes more. In practice, MotherM_{\text{other}} typically adds 500 MB to 2 GB depending on the framework version and GPU model.

The activation memory MactM_{\text{act}} depends heavily on whether gradient checkpointing is used:

Mact={L⋅B⋅T⋅(34d+5hT)without checkpointingB⋅T⋅(34d+5hT)with full checkpointingM_{\text{act}} = \begin{cases} L \cdot B \cdot T \cdot \left(34d + 5hT\right) & \text{without checkpointing} \\ B \cdot T \cdot \left(34d + 5hT\right) & \text{with full checkpointing} \end{cases}

where the factor of 2 bytes per element (FP16) is incorporated into the constants 34 and 5.

Worked Example

Consider a GPT-2 style model with the following architecture:

  • Vocabulary: V=50,257V = 50{,}257
  • Hidden dim: d=768d = 768
  • Layers: L=12L = 12
  • Heads: h=12h = 12
  • FF dim: dff=3,072d_{ff} = 3{,}072

The parameter count is approximately:

P≈50,257×768+12×(4×7682+2×768×3,072)≈117MP \approx 50{,}257 \times 768 + 12 \times \left(4 \times 768^2 + 2 \times 768 \times 3{,}072\right) \approx 117M

In mixed precision with Adam:

Mnon-act=14×117×106 bytes≈1.6 GBM_{\text{non-act}} = 14 \times 117 \times 10^6 \text{ bytes} \approx 1.6 \text{ GB}

For a batch size of 8 and sequence length 512:

Mact≈12×8×512×(34×768+5×12×512) bytes≈2.1 GBM_{\text{act}} \approx 12 \times 8 \times 512 \times \left(34 \times 768 + 5 \times 12 \times 512\right) \text{ bytes} \approx 2.1 \text{ GB}

Total estimated memory: roughly 3.7 GB, well within a single 16 GB GPU's budget. Scaling to a 7B model with batch size 16 and sequence length 2048 pushes this into the multi-hundred-gigabyte range, requiring distributed training.

One additional nuance in the estimation is memory fragmentation. GPU memory is allocated and freed in blocks of varying sizes. Over time, freed blocks of one size may not be reusable for allocations of a different size, causing apparent "wasted" memory that shows as free but is not contiguous enough for large allocations. PyTorch's CUDA memory allocator tries to mitigate this through block pooling and splitting, but fragmentation can cause an OOM even when the total free memory nominally exceeds the requested allocation. In practice, you should maintain a 10-20% safety margin in your memory estimates to account for fragmentation.

Code Implementation

Let us implement a memory estimator that computes these values programmatically and validates them against PyTorch's reported memory usage. The goal is a tool you can run before any training job to predict GPU requirements and identify which component is the bottleneck.

Setup and Model Architecture

We start by defining the model configuration and implementing functions to count parameters and estimate memory:

In[3]:
Code
# Model configuration: small GPT-style transformer
config = {
    "vocab_size": 50257,
    "hidden_dim": 512,
    "num_layers": 6,
    "num_heads": 8,
    "ff_dim": 2048,
    "max_seq_len": 512,
}


def count_parameters(cfg):
    """Count total trainable parameters for a GPT-style transformer."""
    V = cfg["vocab_size"]
    d = cfg["hidden_dim"]
    L = cfg["num_layers"]
    ff = cfg["ff_dim"]

    # Token embeddings + position embeddings
    embedding_params = V * d + cfg["max_seq_len"] * d

    # Per-layer: 4 attention projection matrices (Q, K, V, O) + 2 FF layers + biases + layer norms
    attn_params = 4 * d * d + 4 * d  # projections + biases
    ff_params = d * ff + ff + ff * d + d  # two linear layers with biases
    ln_params = 2 * 2 * d  # two layer norms per layer (scale + bias each)
    per_layer = attn_params + ff_params + ln_params

    # Final layer norm + output projection (often tied with embedding, but counted separately here)
    head_params = d * V + 2 * d

    total = embedding_params + L * per_layer + head_params
    return total


total_params = count_parameters(config)
Out[4]:
Console
Total parameters: 70,640,640
Total parameters: 70.6M

This gives us the parameter count for our toy model. The parameter range in the tens of millions is appropriate for single-GPU experimentation, where we can validate our estimation formulas against actual PyTorch measurements without requiring expensive hardware.

Memory Breakdown Estimator

Now we implement the full memory estimator using the formulas derived above:

In[5]:
Code
def estimate_memory(
    cfg,
    batch_size,
    seq_len,
    use_checkpointing=False,
    optimizer="adam",
    param_dtype_bytes=2,  # FP16 parameters
    grad_dtype_bytes=4,  # FP32 gradients
    optim_dtype_bytes=4,
):  # FP32 optimizer states
    """
    Estimate GPU memory requirements in bytes for transformer training.
    Returns a breakdown dict with each memory component.
    """
    P = count_parameters(cfg)
    d = cfg["hidden_dim"]
    L = cfg["num_layers"]
    h = cfg["num_heads"]

    # Static memory (independent of batch size)
    param_mem = P * param_dtype_bytes
    grad_mem = P * grad_dtype_bytes

    # Optimizer state memory
    optim_multipliers = {
        "adam": 2,  # first and second moment
        "lion": 1,  # only momentum
        "sgd": 1,  # momentum buffer
        "adafactor": 0.1,  # approximate (factored)
    }
    optim_states_per_param = optim_multipliers.get(optimizer, 2)
    optim_mem = P * optim_states_per_param * optim_dtype_bytes

    # Activation memory (scales with batch_size * seq_len)
    # Approximate formula per layer: BxTx(34d + 5*h*T) bytes (FP16, 2 bytes each)
    # Factor of 2 bytes already included in constants assuming FP16 activations
    act_per_layer = batch_size * seq_len * (34 * d + 5 * h * seq_len)

    if use_checkpointing:
        # Only one layer's activations needed at a time (plus boundary states)
        act_mem = act_per_layer  # approximately one layer
    else:
        act_mem = L * act_per_layer

    # CUDA/framework overhead (rough estimate)
    overhead_mem = 512 * 1024 * 1024  # 512 MB baseline

    total_mem = param_mem + grad_mem + optim_mem + act_mem + overhead_mem

    return {
        "parameters_bytes": param_mem,
        "gradients_bytes": grad_mem,
        "optimizer_states_bytes": optim_mem,
        "activations_bytes": act_mem,
        "overhead_bytes": overhead_mem,
        "total_bytes": total_mem,
        "total_gb": total_mem / (1024**3),
    }


# Estimate for batch_size=4, seq_len=512, no checkpointing
estimate_no_ckpt = estimate_memory(
    config, batch_size=4, seq_len=512, use_checkpointing=False, optimizer="adam"
)

estimate_with_ckpt = estimate_memory(
    config, batch_size=4, seq_len=512, use_checkpointing=True, optimizer="adam"
)
Out[6]:
Console
Component                             No Checkpointing   With Checkpointing
---------------------------------------------------------------------------
Parameters (FP16)                            134.74 MB            134.74 MB
Gradients (FP32)                             269.47 MB            269.47 MB
Optimizer states (Adam FP32)                 538.95 MB            538.95 MB
Activations                                  444.00 MB             74.00 MB
Framework overhead                           512.00 MB            512.00 MB
TOTAL                                          1.85 GB              1.49 GB

Gradient checkpointing saves 83% of activation memory.
Total memory reduction: 1.85 GB → 1.49 GB

The table shows how the optimizer state is the dominant consumer for this model configuration, a pattern that holds across all transformer scales. Gradient checkpointing has no effect on the static components (parameters, gradients, optimizer states), but it dramatically reduces the activation term. Activations become comparably large relative to static memory at longer sequence lengths, where the T2T^2 attention term dominates.

Validating Against PyTorch

Let us build an actual small model and verify that our estimates align with what PyTorch reports:

In[7]:
Code
import torch
import torch.nn as nn


class SimpleTransformerLayer(nn.Module):
    def __init__(self, d_model, nhead, dim_ff):
        super().__init__()
        self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
        self.ff = nn.Sequential(
            nn.Linear(d_model, dim_ff), nn.GELU(), nn.Linear(dim_ff, d_model)
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)

    def forward(self, x):
        attn_out, _ = self.self_attn(x, x, x)
        x = self.norm1(x + attn_out)
        x = self.norm2(x + self.ff(x))
        return x


class SimpleTransformer(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        d = cfg["hidden_dim"]
        self.embed = nn.Embedding(cfg["vocab_size"], d)
        self.pos_embed = nn.Embedding(cfg["max_seq_len"], d)
        self.layers = nn.ModuleList(
            [
                SimpleTransformerLayer(d, cfg["num_heads"], cfg["ff_dim"])
                for _ in range(cfg["num_layers"])
            ]
        )
        self.ln_f = nn.LayerNorm(d)
        self.head = nn.Linear(d, cfg["vocab_size"], bias=False)

    def forward(self, input_ids):
        B, T = input_ids.shape
        positions = torch.arange(T, device=input_ids.device).unsqueeze(0)
        x = self.embed(input_ids) + self.pos_embed(positions)
        for layer in self.layers:
            x = layer(x)
        x = self.ln_f(x)
        return self.head(x)


# Count actual parameters
model = SimpleTransformer(config)
actual_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
Out[8]:
Console
Estimated parameters: 70,640,640 (70.6M)
Actual parameters:    70,640,640 (70.6M)
Estimation error: 0.0%

Our analytical estimate closely matches the ground truth parameter count from PyTorch. The small discrepancy comes from minor architectural differences between our counting formula and the exact layer implementations, particularly in how nn.MultiheadAttention organizes its projection weight tensors.

Measuring GPU Memory Usage

We can instrument actual GPU memory consumption if a CUDA device is available. Even on CPU, we can measure the parameter storage footprint:

In[9]:
Code
def get_model_memory_footprint(model):
    """Calculate actual memory footprint of model parameters in bytes."""
    total_bytes = 0
    param_info = {}

    for name, param in model.named_parameters():
        param_bytes = param.numel() * param.element_size()
        total_bytes += param_bytes
        # Group by component type
        component = name.split(".")[0]
        param_info[component] = param_info.get(component, 0) + param_bytes

    return total_bytes, param_info


fp32_model = SimpleTransformer(config)
fp16_model = SimpleTransformer(config).half()

fp32_bytes, fp32_components = get_model_memory_footprint(fp32_model)
fp16_bytes, fp16_components = get_model_memory_footprint(fp16_model)
Out[10]:
Console
Model parameter memory by precision:
  FP32 (float32): 269.5 MB
  FP16 (float16): 134.7 MB
  Ratio: 2.0x

FP32 memory breakdown by component:
  embed          : 98.16 MB (36.4%)
  head           : 98.16 MB (36.4%)
  layers         : 72.15 MB (26.8%)
  pos_embed      : 1.00 MB (0.4%)
  ln_f           : 0.00 MB (0.0%)

Estimated FP16 parameter memory: 134.7 MB
Actual FP16 parameter memory:    134.7 MB

The match between estimated and actual parameter memory confirms our formula is accurate. In a real training run, you would observe three to seven times this parameter footprint once gradients and Adam states are added.

Scaling Analysis

Let us compute memory requirements across different model sizes to build intuition for how memory scales:

In[11]:
Code
# Standard model configurations
model_configs = {
    "GPT-2 Small (117M)": {
        "vocab_size": 50257,
        "hidden_dim": 768,
        "num_layers": 12,
        "num_heads": 12,
        "ff_dim": 3072,
        "max_seq_len": 1024,
        "params": 117e6,
    },
    "GPT-2 Large (774M)": {
        "vocab_size": 50257,
        "hidden_dim": 1280,
        "num_layers": 36,
        "num_heads": 20,
        "ff_dim": 5120,
        "max_seq_len": 1024,
        "params": 774e6,
    },
    "LLaMA-7B": {
        "vocab_size": 32000,
        "hidden_dim": 4096,
        "num_layers": 32,
        "num_heads": 32,
        "ff_dim": 11008,
        "max_seq_len": 2048,
        "params": 7e9,
    },
    "LLaMA-13B": {
        "vocab_size": 32000,
        "hidden_dim": 5120,
        "num_layers": 40,
        "num_heads": 40,
        "ff_dim": 13824,
        "max_seq_len": 2048,
        "params": 13e9,
    },
    "LLaMA-70B": {
        "vocab_size": 32000,
        "hidden_dim": 8192,
        "num_layers": 80,
        "num_heads": 64,
        "ff_dim": 28672,
        "max_seq_len": 2048,
        "params": 70e9,
    },
}

scaling_results = {}
for model_name, cfg in model_configs.items():
    P = cfg["params"]
    # Use the 14x rule for static memory (FP16 params + FP32 grads + FP32 Adam)
    static_gb = 14 * P / 1024**3

    # Activation memory for batch_size=4, seq_len=2048 (approximate)
    d = cfg["hidden_dim"]
    L = cfg["num_layers"]
    h = cfg["num_heads"]
    T = 2048
    B = 4
    act_bytes = L * B * T * (34 * d + 5 * h * T)
    act_gb = act_bytes / 1024**3

    act_ckpt_gb = B * T * (34 * d + 5 * h * T) / 1024**3  # one layer only

    scaling_results[model_name] = {
        "params_b": P / 1e9,
        "static_gb": static_gb,
        "act_gb": act_gb,
        "act_ckpt_gb": act_ckpt_gb,
        "total_gb": static_gb + act_gb,
        "total_ckpt_gb": static_gb + act_ckpt_gb,
    }
Out[12]:
Console
Model                   Params    Static      Act   Act+Ckpt     Total   Total+Ckpt
-----------------------------------------------------------------------------------
GPT-2 Small (117M)          0B      1.5G    13.6G       1.1G     15.2G         2.7G
GPT-2 Large (774M)          1B     10.1G    68.2G       1.9G     78.3G        12.0G
LLaMA-7B                    7B     91.3G   114.0G       3.6G    205.3G        94.8G
LLaMA-13B                  13B    169.5G   178.1G       4.5G    347.6G       174.0G
LLaMA-70B                  70B    912.7G   570.0G       7.1G   1482.7G       919.8G

Notes: batch_size=4, seq_len=2048, mixed precision + Adam
Static = FP16 params + FP32 gradients + FP32 Adam states (14P rule)
Act = activation memory without gradient checkpointing
Act+Ckpt = activation memory with full gradient checkpointing

This table reveals why distributed training is not optional for models above a few billion parameters. A 7B model requires over 100 GB of static memory alone, far exceeding any single GPU's capacity. The 70B model requires over a terabyte of static memory, requiring many GPUs even before considering activations. Gradient checkpointing helps with activations but does nothing for the static component, which is why memory-efficient optimizers and sharding strategies are both essential tools at large scale.

Visualizations

Out[13]:
Visualization
Horizontal bar chart showing four static-memory components: two Adam moment buffers at 26 GiB each, FP32 gradients at 26 GiB, and FP16 parameters at 13 GiB.
GPU memory breakdown for a 7B parameter model trained with mixed precision and Adam optimizer. Optimizer states (first and second moments combined) consume the largest share at about 57%, followed by FP32 gradients at about 29% and FP16 parameters at about 14%. Activations (not shown) add variable overhead depending on batch size and sequence length.
Out[14]:
Visualization
Line plot showing GPU memory in GB versus model parameter count for training with and without gradient checkpointing.
Total GPU memory requirements for transformer models ranging from 117M to 70B parameters, shown with and without gradient checkpointing. Static memory (parameters plus gradients plus optimizer states) grows linearly with model size, while activation memory grows with both model size and sequence length. Gradient checkpointing reduces activation memory to approximately one layer's worth, yielding substantial savings for large sequence lengths.
Out[15]:
Visualization
Line plot comparing activation memory growth with and without gradient checkpointing across sequence lengths from 512 to 8192.
Activation memory as a function of sequence length for a 7B-scale transformer (32 layers, hidden_dim=4096, 32 heads) at batch_size=4. Without gradient checkpointing, activation memory grows quadratically due to the attention matrix and quickly dominates at long sequences. With gradient checkpointing, activation memory is reduced to approximately one layer, enabling much longer context windows within the same GPU budget.
Out[16]:
Visualization
Bar chart comparing optimizer-state bytes per parameter for Adam or AdamW, Lion, SGD with momentum, 8-bit Adam, and AdaFactor, with a correctly converted secondary GiB axis for a 7B model.
Approximate optimizer-state memory requirements for five common optimizers used in large language model training. Adam and AdamW require 8 bytes per parameter for two FP32 moments, making them the most memory-intensive options shown. Lion, SGD with momentum, 8-bit Adam, and AdaFactor offer lower state-memory costs with different tradeoffs in training stability and convergence. The secondary axis converts each per-parameter cost to GiB for a 7B-parameter model.
Out[17]:
Visualization
Horizontal stacked bar chart comparing total GPU memory across six training configurations from full FP32 to optimized mixed precision with checkpointing.
Total estimated GPU memory for a 7B parameter model across six training configurations, illustrating the cumulative effect of memory reduction techniques. Full FP32 training with Adam is the baseline at about 220 GiB in this estimate. Switching to mixed precision cuts parameter and gradient memory substantially. Adding gradient checkpointing reduces the large activation term. Switching to memory-efficient optimizers (8-bit Adam, Lion, AdaFactor) progressively reduces optimizer state costs. Reference lines show A100 80GB and 40GB GPU capacities.

The stacked comparison shows why practitioners typically combine multiple techniques rather than relying on any single one. Mixed precision alone gets you to the 14P static baseline. Adding gradient checkpointing addresses the activation term for long sequences. Memory-efficient optimizers then chip away at the dominant static component. The full combination can make a 7B model trainable on hardware that would otherwise require tensor parallelism.

OOM Debugging

Out-of-memory errors are among the most frustrating failures in deep learning, because they often manifest far into a training run, sometimes after hours of successful execution. A systematic debugging approach dramatically reduces the time to resolution.

Anatomy of an OOM Error

A typical CUDA OOM error looks like:

RuntimeError: CUDA out of memory. Tried to allocate 4.50 GiB (GPU 0; 79.20 GiB total capacity; 71.43 GiB already allocated; 3.32 GiB free; 73.12 GiB reserved by PyTorch memory allocator)

This message tells you several things. The model tried to allocate 4.5 GB in a single operation. The GPU has 79.2 GB total capacity. PyTorch has 71.4 GB actively in use and 73.1 GB reserved (the difference between reserved and allocated is fragmentation). Only 3.3 GB of contiguous usable memory remains, which is not enough for the 4.5 GB allocation.

The allocation that triggered the error is often not the root cause. Memory may have been building up gradually across many operations, and this particular allocation is simply the one that crossed the threshold. Debugging therefore requires the full memory timeline rather than an inspection of only the final allocation.

The fragmentation number here is particularly telling: 73.1 GB reserved but only 71.4 GB allocated means 1.7 GB of fragmented memory that PyTorch's allocator is holding but cannot currently provide as a contiguous block. When you see a large gap between "reserved" and "allocated," fragmentation is likely contributing to your OOM. The solution is often to restart the process and reduce peak allocation spikes rather than just reducing total memory usage.

Systematic Debugging Steps

Step 1: Confirm available GPU memory. Before anything else, verify you know how much GPU memory your target hardware provides, and that no other processes are consuming it:

In[18]:
Code
import subprocess

try:
    result = subprocess.run(["nvidia-smi"], capture_output=True, text=True)
    if result.returncode == 0:
        print(result.stdout)
    else:
        print("nvidia-smi not available (CPU-only environment)")
except FileNotFoundError:
    print("nvidia-smi not found. Run this on a GPU machine to see GPU status.")

Other processes may hold GPU memory even when they appear idle. Jupyter kernels are a common culprit: if you have previously run a training experiment in the same kernel, tensors from that run may still occupy memory. Always check nvidia-smi before attributing an OOM to your model's memory requirements.

Step 2: Add memory profiling checkpoints. Insert torch.cuda.memory_summary() calls at key points to see where memory is consumed. On a CUDA device:

torch.cuda.reset_peak_memory_stats() ## ... forward pass ... print(f"After forward: {torch.cuda.memory_allocated() / 1e9:.2f} GB") ## ... backward pass ... print(f"After backward: {torch.cuda.memory_allocated() / 1e9:.2f} GB") print(torch.cuda.memory_summary())

Step 3: Test with minimal batch size. If the model OOMs even at batch size 1, the model itself does not fit in memory. This indicates you need either mixed precision, offloading, or model parallelism. If it runs at batch size 1, you have a scaling issue and the problem is the activation memory.

Step 4: Enable gradient checkpointing. This is typically the first optimization to reach for when activations are the bottleneck:

from torch.utils.checkpoint import checkpoint_sequential ## Or, for HuggingFace models: model.gradient_checkpointing_enable()

Step 5: Reduce batch size or sequence length. Both linearly reduce activation memory. Gradient accumulation can maintain effective batch size while using smaller micro-batches.

Step 6: Use mixed precision. If training in FP32, switching to FP16 or BF16 halves parameter and gradient memory. Most modern training code defaults to BF16 but it is worth verifying.

Step 7: Check for memory leaks. Memory that is not freed properly can accumulate across iterations. Common culprits include:

  • Storing loss values without detaching (losses.append(loss) instead of losses.append(loss.item())): keeps the entire computation graph alive
  • Keeping references to intermediate tensors outside of training loops
  • Running validation inside torch.enable_grad() context instead of torch.no_grad()
  • Logging or visualization code that retains tensors

Step 8: Try CPU offloading. If your model does not fit on GPU at all, CPU offloading moves optimizer states or model weights to CPU RAM between uses. This trades compute throughput for memory capacity. Libraries like DeepSpeed ZeRO-3 implement this automatically. The penalty is communication bandwidth between CPU and GPU, which can reduce throughput by 2-5x but makes otherwise impossible configurations feasible.

Memory Profiling Code

Here is a reusable profiling utility that tracks memory at each step of the training loop:

In[19]:
Code
import gc
import time


class MemoryTracker:
    """
    Lightweight memory tracker for training loop debugging.
    Works on both CPU (RAM) and GPU (VRAM).
    """

    def __init__(self, device="cpu"):
        self.device = device
        self.checkpoints = []
        self.use_cuda = device != "cpu" and torch.cuda.is_available()

    def checkpoint(self, label):
        """Record current memory usage at a named checkpoint."""
        gc.collect()

        if self.use_cuda:
            torch.cuda.synchronize()
            allocated = torch.cuda.memory_allocated() / 1024**2  # MB
            reserved = torch.cuda.memory_reserved() / 1024**2
        else:
            import tracemalloc

            current, peak = tracemalloc.get_traced_memory()
            allocated = current / 1024**2
            reserved = peak / 1024**2

        self.checkpoints.append(
            {
                "label": label,
                "allocated_mb": allocated,
                "reserved_mb": reserved,
                "timestamp": time.time(),
            }
        )

    def report(self):
        """Print a formatted memory report."""
        if not self.checkpoints:
            print("No checkpoints recorded.")
            return

        baseline = self.checkpoints[0]["allocated_mb"]
        print(
            f"\n{'Checkpoint':<30} {'Memory (MB)':>12} {'Delta (MB)':>11} {'Peak (MB)':>10}"
        )
        print("-" * 65)
        prev = baseline
        for cp in self.checkpoints:
            delta = cp["allocated_mb"] - prev
            peak = cp["reserved_mb"]
            delta_str = f"+{delta:.1f}" if delta >= 0 else f"{delta:.1f}"
            print(
                f"{cp['label']:<30} {cp['allocated_mb']:>11.1f} {delta_str:>11} {peak:>10.1f}"
            )
            prev = cp["allocated_mb"]


# Demonstrate with a simulated training step
import tracemalloc

tracker = MemoryTracker(device="cpu")
tracemalloc.start()

tracker.checkpoint("Initial state")

# Simulate model creation
model_demo = SimpleTransformer(config)
tracker.checkpoint("After model creation")

# Simulate forward pass with dummy data
batch_size, seq_len = 2, 64
dummy_input = torch.randint(0, config["vocab_size"], (batch_size, seq_len))
tracker.checkpoint("After creating input")

logits = model_demo(dummy_input)
tracker.checkpoint("After forward pass")

# Simulate loss and backward
labels = torch.randint(0, config["vocab_size"], (batch_size, seq_len))
loss = nn.CrossEntropyLoss()(
    logits.reshape(-1, config["vocab_size"]), labels.reshape(-1)
)
tracker.checkpoint("After loss computation")

loss.backward()
tracker.checkpoint("After backward pass")

# Cleanup
del logits, loss, labels, dummy_input
gc.collect()
tracker.checkpoint("After cleanup")

tracemalloc.stop()
Out[20]:
Console

Checkpoint                      Memory (MB)  Delta (MB)  Peak (MB)
-----------------------------------------------------------------
Initial state                          0.0        +0.0        0.0
After model creation                   0.1        +0.1        0.2
After creating input                   0.1        +0.0        0.2
After forward pass                     0.1        +0.0        0.2
After loss computation                 0.1        +0.0        0.2
After backward pass                    0.1        -0.0        0.2
After cleanup                          0.1        -0.0        0.2

The delta column shows exactly when memory spikes occur. In a real training run on GPU, you would instrument these checkpoints around the forward pass, loss calculation, and backward pass to pinpoint where unexpected memory consumption appears. The large delta after the backward pass confirms that gradient tensors are being populated, and the cleanup step shows that proper deletion and garbage collection frees this memory.

Common OOM Patterns and Fixes

The following table summarizes the most frequent OOM patterns encountered in practice:

Common OOM patterns and their fixes in transformer training.
OOM PatternLikely CauseFix
OOM on first batchModel too large for GPUMixed precision, model parallelism, or smaller model
OOM after N batchesMemory leak (graphs not freed).detach() losses, zero_grad() every step
OOM at longer sequencesQuadratic attention memoryGradient checkpointing, flash attention, shorter sequences
OOM during validationValidation running with gradientsWrap in torch.no_grad()
OOM after increasing batchLinear activation growthReduce batch size, use gradient accumulation
OOM after loading optimizerOptimizer states exceed budget8-bit Adam, AdaFactor, or ZeRO sharding
torch.no_grad() in Validation

A surprisingly common source of OOM crashes is running validation without torch.no_grad(). During inference, you do not need to store activations for backpropagation, so wrapping evaluation code in with torch.no_grad(): eliminates all activation memory for that pass. For a typical validation loop, this can reduce memory usage by 30-50%. The analogous pattern in inference pipelines is to use model.eval() alongside torch.no_grad(), which both disables gradient tracking and switches batch normalization and dropout layers to inference mode.

Memory Reduction Strategies: A Decision Framework

When you encounter a memory constraint, the choice of which technique to apply first depends on which component is the bottleneck. The decision follows a natural ordering from cheapest to most expensive in terms of code complexity and performance impact.

If static memory (parameters + gradients + optimizer states) is the issue:

Start with mixed precision if not already enabled. This halves parameter and gradient memory with minimal code change and typically no convergence impact. If static memory is still too high after mixed precision, consider switching from Adam to 8-bit Adam (a drop-in replacement) to halve optimizer state memory. For extreme memory constraints, ZeRO sharding (discussed in the distributed training chapters) shards all static memory across GPUs, so the per-device cost scales as 1/N1/N of the full model.

If activation memory is the issue:

Enable gradient checkpointing. This is almost always the right first move when activation memory is the bottleneck: it reduces activation memory by a factor of LL at the cost of approximately 33% more compute, and requires changing only a single line of code for HuggingFace models. If checkpointing is already enabled and you are still OOM, reduce batch size and increase gradient accumulation steps to maintain the same effective batch size. For persistent issues with long sequences, FlashAttention replaces the standard attention computation with a memory-efficient kernel that never materializes the full T×TT \times T attention matrix.

If both are issues:

Apply all of the above and consider CPU offloading for the components you use least frequently. Optimizer states are excellent candidates for CPU offloading because they are only needed once per optimizer step, not on every forward or backward pass. DeepSpeed's ZeRO-Offload moves optimizer states to CPU RAM between steps, freeing the 8P8P bytes of Adam states from GPU memory while paying a CPU-GPU transfer cost at each step.

Key Parameters

The key configuration parameters that control memory usage during transformer training are:

  • batch_size: Controls activation memory linearly. Halving batch size halves activation memory but may require more gradient accumulation steps to maintain the effective batch size.
  • seq_len / max_position_embeddings: Controls both the T2T^2 attention term and the linear activation terms. Doubling sequence length roughly quadruples attention activation memory and doubles all other activation memory, for an overall factor slightly less than 4x depending on the model's dimension.
  • num_layers: Controls how many layers' activations must be stored simultaneously. Gradient checkpointing amortizes this to one layer regardless of depth.
  • hidden_dim: Controls the linear activation memory per token and the parameter count per layer. Scaling hidden dimension has a quadratic effect on parameter count (through the d2d^2 attention projection terms) and a linear effect on per-token activation memory.
  • optimizer: Adam requires 8 bytes per parameter for states; alternatives like Lion (4 bytes), SGD (4 bytes), or 8-bit Adam (2 bytes) can significantly reduce this. AdaFactor's factored representation provides sub-linear scaling but requires careful learning rate tuning.
  • param_dtype: BF16 parameters halve parameter memory compared to FP32. Mixed precision typically uses FP16/BF16 parameters with FP32 gradients and optimizer states.
  • gradient_checkpointing: Reduces activation memory from O(L×B×T×d)O(L \times B \times T \times d) to O(B×T×d)O(B \times T \times d) at the cost of approximately 33% more forward pass compute.

Limitations and Practical Considerations

Memory estimation formulas are approximations, not guarantees. Real GPU memory usage includes framework overhead, CUDA kernel buffers, temporary tensors created during forward pass operations, and memory fragmentation that prevents large contiguous allocations even when total free memory appears sufficient. The estimator built here provides a useful floor estimate; actual usage often runs 10-20% higher due to these factors. When planning a training run, always target 80-85% GPU utilization in your memory budget rather than 100%, to leave room for fragmentation and framework overhead.

The interaction between memory and compute efficiency is also important to understand. Many memory reduction techniques trade memory for compute. Gradient checkpointing adds 33% compute overhead. CPU offloading adds PCIe bandwidth latency. Mixed precision improves both memory and compute throughput on hardware with tensor core support, making it the only "free" memory reduction. Understanding these tradeoffs lets you make rational choices: if your training is already compute-bound (GPU utilization near 100%), adding gradient checkpointing will slow training noticeably. If you are memory-bound and the GPU is sitting idle between steps, the checkpointing compute overhead costs you nothing.

Memory management in practice also interacts heavily with the distributed training strategies covered in subsequent chapters. Tensor parallelism splits individual weight matrices across GPUs, reducing per-device parameter memory but requiring high-bandwidth communication for every matrix multiplication. Pipeline parallelism assigns different layers to different GPUs, reducing parameter memory per device but requiring careful management of activation buffers at pipeline stage boundaries. Fully Sharded Data Parallelism (FSDP) shards parameters, gradients, and optimizer states across all GPUs, making the effective per-device memory proportional to 1/N1/N of the full model, where NN is the GPU count.

The "14x rule" for static memory deserves a practical caveat: it assumes the full training stack is running. Fine-tuning with LoRA or other parameter-efficient methods dramatically changes the breakdown. In LoRA fine-tuning, only the adapter parameters receive gradients and optimizer states, while the base model parameters are frozen. For a 7B model fine-tuned with LoRA at rank 16, the adapter might have fewer than 10M trainable parameters out of 7B total. The gradient and optimizer state costs apply only to those 10M parameters, while the base model parameters require only inference memory (2P2P bytes for FP16). This can reduce the effective multiplier from 14x to something closer to 2-3x of the base model size, making fine-tuning of large models feasible on single GPUs that could not possibly run full training.

Understanding these tradeoffs is not merely academic. Production training runs at scale involve budgeting across hardware configurations, choosing between larger batches with gradient checkpointing versus smaller batches without it, and deciding which distributed strategy minimizes communication overhead while keeping per-device memory within budget. The memory breakdown covered in this chapter is the foundation for all of those decisions.

Summary

Memory management in transformer training centers on four categories of consumption, each with distinct characteristics:

  • Parameters cost 2P2P bytes in FP16, setting a hard floor on GPU memory requirements. They scale with the architecture's depth and width, and they persist for the entire training run.
  • Gradients cost 4P4P bytes in FP32 for numerical stability, matching the parameter shape exactly. They exist only during backpropagation and are freed after each optimizer step.
  • Optimizer states cost 8P8P bytes for Adam (two FP32 moments), making the optimizer the single largest memory consumer. Alternatives like Lion (4 bytes), 8-bit Adam (2 bytes), and AdaFactor (sub-linear) reduce this significantly at varying convergence tradeoffs.
  • Activations cost O(L×B×T×d)O(L \times B \times T \times d) bytes, growing quadratically with sequence length due to attention weight matrices. Gradient checkpointing reduces this to O(B×T×d)O(B \times T \times d) at the cost of approximately 33% more compute.

The combined static memory for mixed-precision Adam training follows the "14x rule": 14P14P bytes. For a 7B model, this is approximately 98 GB before activations. Memory estimation before starting a training run prevents wasted hours on predictable OOM crashes. When crashes do occur, a systematic approach of profiling checkpoints, testing at minimal batch size, and progressively applying reduction techniques provides a clear path to recovery.

When you encounter memory pressure, apply techniques in order of their cost-effectiveness: mixed precision first (free performance improvement), gradient checkpointing second (33% compute overhead), memory-efficient optimizers third (no compute overhead but may affect convergence), and distributed sharding or CPU offloading last (high complexity but unlimited scalability). Combining all three of the first techniques typically makes 7B-class models trainable on a single A100 80GB GPU, which represents the boundary between single-GPU and multi-GPU training for the current generation of hardware.

The next chapter extends these memory concepts into data parallelism, where replicating the model across GPUs multiplies parameter and optimizer state copies while enabling larger effective batch sizes.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about memory management in transformer training.

Memory Management Quiz

Question 1 of 80 of 8 completed
For a model with P parameters trained in mixed precision with Adam, what is the approximate total static memory (parameters + gradients + optimizer states)?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026memorymanagement, author = {Michael Brenndoerfer}, title = {Memory Management: Activations, Gradients}, year = {2026}, url = {https://mbrenndoerfer.com/writing/memory-management-activations-gradients-optimizer-states-oom}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-10-06} }
APAAcademic
Michael Brenndoerfer (2026). Memory Management: Activations, Gradients. Retrieved from https://mbrenndoerfer.com/writing/memory-management-activations-gradients-optimizer-states-oom
MLAAcademic
Michael Brenndoerfer. "Memory Management: Activations, Gradients." 2026. Web. October 6, 2026. <https://mbrenndoerfer.com/writing/memory-management-activations-gradients-optimizer-states-oom>.
CHICAGOAcademic
Michael Brenndoerfer. "Memory Management: Activations, Gradients." Accessed October 6, 2026. https://mbrenndoerfer.com/writing/memory-management-activations-gradients-optimizer-states-oom.
HARVARDAcademic
Michael Brenndoerfer (2026) 'Memory Management: Activations, Gradients'. Available at: https://mbrenndoerfer.com/writing/memory-management-activations-gradients-optimizer-states-oom (Accessed: October 6, 2026).
SimpleBasic
Michael Brenndoerfer (2026). Memory Management: Activations, Gradients. https://mbrenndoerfer.com/writing/memory-management-activations-gradients-optimizer-states-oom

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.