GPU memory, gradient accumulation and batches
Be able to compute the memory requirement and use gradient accumulation and checkpointing.
Prerequisites
- DTrain a neural network in PyTorchrequired
- EParallelism and why GPUsrequired
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 item | The size | Depends on |
|---|---|---|
| The parameters | bytes | the model |
| The gradients | bytes | the model |
| The optimiser state | 4 B for Adam | the model |
| The activations | ∝ the batch size × the sequence length × the depth | the 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 item | The 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 activations | 12 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 of them.
Mathematically it is nearly identical to a large batch. Two details:
- Divide the loss by , otherwise the gradient becomes times too large.
- 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 memory | The time | |
|---|---|---|
| Without | 1× | |
| With, every layer | ~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 measure | The cost |
|---|---|---|
| 1 | Reduce the microbatch, increase the accumulation | none (nearly) |
| 2 | bf16/fp16 instead of fp32 | no loss of quality in practice |
| 3 | Activation checkpointing | ~30 % longer |
| 4 | An 8-bit optimiser | a small quality risk |
| 5 | Shorter sequences | depends on the task |
| 6 | Sharding (ZeRO/FSDP) over several GPUs | requires more cards |
| 7 | A smaller model | the last resort |
Two common mistakes that give an OOM without it being the model's fault:
- Accumulating tensors with their gradient graph.
total_loss += losssaves the whole graph for every batch. Useloss.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
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- arXiv — Training Deep Nets with Sublinear Memory Cost — arXiv (open access; licence per article)
- arXiv — ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — arXiv (open access; licence per article)