Skip to content
AI-grafen
EUniversityModel training and fine-tuning· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

GPU memory, gradient accumulation and batches

Be able to compute the memory requirement and use gradient accumulation and checkpointing.

Prerequisites

Intuition

«CUDA out of memory» is the most common error in deep learning. What will fit can be worked out in advance.

Four items take memory during training:

The itemThe sizeDepends on
The parametersP×P \times bytesthe model
The gradientsP×P \times bytesthe model
The optimiser state2P×2P \times 4 B for Adamthe model
The activations∝ the batch size × the sequence length × the depththe batch

The first three are fixed — they depend only on the model. The fourth grows with the batch, which is why «reduce the batch size» is the first thing you do.

An example, a model with 1 billion parameters in bf16 with Adam:

The itemThe memory
The parameters (2 B)2 GB
The gradients (2 B)2 GB
Adam's m and v (4 B each)8 GB
The total before the activations12 GB

On a 24 GB card about 12 GB therefore remains for the activations — and they are what decides how large a batch you can run.

Formal

Gradient accumulation gives a large effective batch without large memory: run several small batches, sum the gradients, and take a step only after kk of them.

Beffective=Bmicro×k×the number of GPUsB_{\text{effective}} = B_{\text{micro}} \times k \times \text{the number of GPUs}

Mathematically it is nearly identical to a large batch. Two details:

  1. Divide the loss by kk, otherwise the gradient becomes kk times too large.
  2. BatchNorm normalises per microbatch, so the statistics differ from a genuine large batch. LayerNorm and GroupNorm are unaffected — one more reason why transformers use LayerNorm.

Activation checkpointing trades computation for memory: save the activations only at certain points and recompute the rest in the backward pass.

The memoryThe time
WithoutO(L)O(L)1×
With, every layerO(L)O(\sqrt{L})~1.3×

Thirty per cent more training time in order to be able to run a model that otherwise does not fit is nearly always a good deal.

The list of measures at an OOM, in order of how much they cost:

#The measureThe cost
1Reduce the microbatch, increase the accumulationnone (nearly)
2bf16/fp16 instead of fp32no loss of quality in practice
3Activation checkpointing~30 % longer
4An 8-bit optimisera small quality risk
5Shorter sequencesdepends on the task
6Sharding (ZeRO/FSDP) over several GPUsrequires more cards
7A smaller modelthe last resort

Two common mistakes that give an OOM without it being the model's fault:

  • Accumulating tensors with their gradient graph. total_loss += loss saves the whole graph for every batch. Use loss.item() or .detach().
  • Forgetting torch.no_grad() at evaluation. Then a graph that is never used is built, and the memory doubles.

Code

import torch, torch.nn as nn

def memory_estimate(parameters, bytes_per_param=2, optimiser="adam"):
    """GB for the parameters, the gradients and the optimiser state — everything but the activations."""
    per_opt = {"adam": 8, "adamw": 8, "sgd_momentum": 4, "sgd": 0, "adam8bit": 2}[optimiser]
    total_bytes = parameters * (bytes_per_param * 2 + per_opt)
    return round(total_bytes / 1024**3, 2)

for n, name in ((125e6, "125M"), (1e9, "1B"), (7e9, "7B")):
    print(f"{name:>5}: adam {memory_estimate(n):>6.2f} GB   "
          f"adam8bit {memory_estimate(n, optimiser='adam8bit'):>6.2f} GB")
#  125M:   1.40 GB     0.70 GB
#    1B:  11.18 GB     5.59 GB
#    7B:  78.23 GB    39.12 GB

# Gradient accumulation
ACC = 8                                      # the effective batch = the microbatch × 8
opt.zero_grad()
for i, (x, y) in enumerate(dataloader):
    loss = nn.functional.cross_entropy(model(x), y) / ACC       # ← divide!
    loss.backward()
    if (i + 1) % ACC == 0:
        nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        opt.step()
        opt.zero_grad()

# Activation checkpointing
from torch.utils.checkpoint import checkpoint

class Stack(nn.Module):
    def __init__(self, blocks, use_checkpoint=True):
        super().__init__()
        self.blocks = nn.ModuleList(blocks)
        self.cp = use_checkpoint

    def forward(self, x):
        for b in self.blocks:
            x = checkpoint(b, x, use_reentrant=False) if self.cp and self.training else b(x)
        return x

# Measure the actual peak memory
def peak_memory(fn):
    torch.cuda.reset_peak_memory_stats()
    fn()
    return round(torch.cuda.max_memory_allocated() / 1024**3, 2)

# The two classic mistakes
total = 0.0
for x, y in dataloader:
    loss = nn.functional.cross_entropy(model(x), y)
    total += loss.item()        # RIGHT — .item() breaks the graph
    # total += loss             # WRONG — saves the whole graph, the memory grows every batch

model.eval()
with torch.no_grad():           # RIGHT — without this a graph that is never used is built
    for x, y in val_loader:
        model(x)

Mastery means

  • Computes the memory requirement of a training run
  • Uses gradient accumulation
  • Knows what activation checkpointing costs and gives

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences