GPU Memory and Speed

What is using your GPU memory

Training memory splits into four tenants — weights, gradients, optimizer state and activations — and you can predict each one with multiplication before ever touching a GPU.

On this page 5
  1. Why you should care
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Four things share your GPU's memory during training: the model's weights, their corrections, the optimizer's notebooks, and the rough work for the current batch.

Picture a student's desk during an exam. The textbook takes fixed space. So does the answer sheet, and the teacher's record register. But the biggest mess is the rough work — pages of in-between scribbles that pile up with every question attempted.

GPU memory during training looks exactly like that desk. And the rough work — not the textbook — is usually what fills it.

Why you should care

When memory runs out, people's first instinct is "my model is too big". Often wrong. The four tenants grow for different reasons, so the right fix depends on which tenant is the fat one.

  • Weights — the model's learned numbers. Fixed size, whatever the batch.
  • Gradients — one correction slot per weight. Same size as the weights.
  • Optimizer state — the popular optimizer, Adam, keeps two extra notebooks per weight.
  • Activations — every in-between result, kept so corrections can be computed later. This one grows with batch size.

How it works

GPU memory during one training step

[ weights ][ gradients ][ optimizer state ][ activations....... ]
   fixed       fixed         fixed            grows with batch
                                              grows with layers

Batch of 8 too big? The activations are the problem. Model will not even load? The fixed tenants are the problem. Different villains, different fixes.

A real example you have seen

This is why a phone can run a photo filter model but not train one. Running a model needs only the weights and a little scratch space. Training needs all four tenants — several times the memory, on the same model.

Remember this

  • Training memory has four tenants, and only activations depend on batch size.
  • Adam quietly quadruples the fixed cost: weights, gradients, two notebooks.
  • Identify the fat tenant first — each one has a different fix.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

The prediction runs anywhere, no GPU needed. The measurement at the end was captured with torch 2.5.1 on an NVIDIA RTX A6000; your exact megabytes will vary a little with version and card.

Predict it with multiplication

anatomy.py
import torch
import torch.nn as nn

model = nn.Sequential(
    nn.Linear(1024, 4096), nn.ReLU(),
    nn.Linear(4096, 4096), nn.ReLU(),
    nn.Linear(4096, 1024),
)
n = sum(p.numel() for p in model.parameters())
mb = n * 4 / 1024**2          # float32 = 4 bytes per number

print(f"parameters:        {n:,} numbers = {mb:6.1f} MB")
print(f"gradients:         one more copy   = {mb:6.1f} MB")
print(f"Adam exp_avg:      one more copy   = {mb:6.1f} MB")
print(f"Adam exp_avg_sq:   one more copy   = {mb:6.1f} MB")
print(f"total before any data arrives:      {4 * mb:6.1f} MB")

opt = torch.optim.Adam(model.parameters())
x = torch.randn(32, 1024)
loss = model(x).sum()
loss.backward()
opt.step()

state = opt.state[next(iter(model.parameters()))]
print("optimizer keeps per-parameter:", sorted(state.keys()))
Output
parameters:        25,175,040 numbers =   96.0 MB
gradients:         one more copy   =   96.0 MB
Adam exp_avg:      one more copy   =   96.0 MB
Adam exp_avg_sq:   one more copy   =   96.0 MB
total before any data arrives:       384.1 MB
optimizer keeps per-parameter: ['exp_avg', 'exp_avg_sq', 'step']

A 96 MB model costs 384 MB before the first sample arrives. That factor of four is Adam's rent: exp_avg (the running average of gradients) and exp_avg_sq (their running square) are full-size tensors, one pair per parameter. See buffers vs parameters for where such state lives.

Now watch it happen live

anatomy_gpu.py
import torch
import torch.nn as nn

if not torch.cuda.is_available():
    raise SystemExit("needs a GPU")

def used():
    return torch.cuda.memory_allocated() / 1024**2

model = nn.Sequential(
    nn.Linear(1024, 4096), nn.ReLU(),
    nn.Linear(4096, 4096), nn.ReLU(),
    nn.Linear(4096, 1024),
).cuda()
print(f"after model.cuda():      {used():7.1f} MB")

opt = torch.optim.Adam(model.parameters())
x = torch.randn(2048, 1024, device="cuda")
loss = model(x).pow(2).mean()
print(f"after forward:           {used():7.1f} MB")

loss.backward()
print(f"after backward:          {used():7.1f} MB")

opt.step()
print(f"after first opt.step():  {used():7.1f} MB")
print(f"peak during all of this: {torch.cuda.max_memory_allocated() / 1024**2:7.1f} MB")
Output
after model.cuda():         96.0 MB
after forward:             184.2 MB
after backward:            216.3 MB
after first opt.step():    408.4 MB
peak during all of this:   504.4 MB

Read it line by line. The forward pass added ~89 MB of activations — outputs of every layer, kept for the backward pass (the mechanics are in autograd graph mechanics). Backward added the gradients and released most activations. The first opt.step() created Adam's two notebooks — the biggest single jump. The peak exceeds every snapshot: mid-backward, activations and gradients overlap.

Common mistakes

Judging memory with nvidia-smi. It reports what the caching allocator reserved from the driver, not your live tensors. Use torch.cuda.memory_allocated() for tensors, memory_reserved() for the allocator, and the gap between them is cache.

Forgetting the optimizer exists. Doubling model width and watching memory quadruple feels like a bug. It is Adam's rent. Plain SGD without momentum carries no notebooks — a real option when memory is desperate.

Blaming parameters for a batch-size crash. If batch 4 fits and batch 8 does not, the fixed tenants are irrelevant. Attack activations: mixed precision, gradient checkpointing, or a smaller batch with accumulation.

Measuring without resetting the peak. max_memory_allocated() remembers the all-time high. Call torch.cuda.reset_peak_memory_stats() before the region you want to measure.

Try it yourself

Swap torch.optim.Adam for torch.optim.SGD(model.parameters(), lr=0.01) and rerun the live script. Predict the "after first opt.step()" line before looking. Then try SGD with momentum=0.9 and explain the difference.

What to learn next

Researcher — Mathematics and papers.

The full accounting

For a model with $P$ parameters trained in float32 with Adam:

$$ M_{\text{fixed}} = \underbrace{4P}{\text{weights}} + \underbrace{4P}{\text{grads}} + \underbrace{8P}_{m, v} = 16P \ \text{bytes} $$

Mixed precision with float16 compute keeps a float32 master copy, shifting the split to $2P + 2P + 4P + 8P = 16P$ — memory-neutral on the fixed side; the saving is in activations. This is the arithmetic behind ZeRO/FSDP: $16P$ divided across $N$ devices.

Activation memory for a transformer block, per token, in bytes (Korthikanti et al., 2022, without selective recomputation):

$$ A \approx s b h \left(34 + 5\frac{a s}{h}\right) $$

  • $s$ — sequence length, $b$ — micro-batch size, $h$ — hidden dimension, $a$ — attention heads.
  • The $\frac{as}{h}$ term is the attention score matrix; FlashAttention eliminates it, and gradient checkpointing divides the rest by trading recomputation.

Why the peak is not the sum

Allocation lifetimes overlap partially. Gradients materialise while activations are still being consumed, so the true peak sits between $\max$ and $\sum$ of the tenants and depends on execution order. torch.cuda.memory._record_memory_history() captures per-allocation stack traces; the flame-graph view at pytorch.org/memory_viz resolves peak attribution exactly.

References

  • Rajbhandari et al. (2020), ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — the $16P$ decomposition and its sharding.
  • Korthikanti et al. (2022), Reducing Activation Recomputation in Large Transformer Models — the activation formula above and selective recomputation.
  • Kingma and Ba (2015), Adam: A Method for Stochastic Optimization — where $m$ and $v$ come from.

What to learn next