skip to content
Victor Guerra

Notes / Training Dynamics & Optimization

Memory Management During Training

Updated Aug 19, 20262 min read
Table of Contents

Where GPU memory goes during training, and the graph-lifetime pitfall that silently OOMs you.


The major consumers

For a model with N parameters:

ConsumerSizeNotes
ParametersN × (4 bytes fp32 / 2 bytes fp16/bf16)the weights themselves
Gradientssame as parameters (N).grad for every param → roughly doubles param memory
Optimizer stateAdam:N (1st + 2nd moment); SGD+momentum:NAdam ⇒ params + grads + 2 moments ≈ 4× param memory (the “×3 extra” over the weights)
Activationsscales with batch × depth × width × seq_lenintermediates cached for backward (see autograd-and-autodiff); often the dominant term
Graph metadataproportional to live activationswhile the graph exists, all its intermediate tensors stay pinned

Rule of thumb (Adam, fp32): parameters + gradients + optimizer state ≈ 16 bytes/param before activations even enter — 4 (param) + 4 (grad) + 8 (Adam moments). Activations are then added on top and are what actually scales with batch size.


The graph-lifetime pitfall — accumulating loss tensors

The bug: summing loss tensors across batches without extracting the scalar:

total_loss += loss # ⚠️ loss still carries grad_fn → keeps its WHOLE graph alive

Because loss is a graph node, holding a reference keeps its entire computation graph (and all cached activations) from being freed. Over many batches these chain together → memory grows every step → OOM that appears partway through the epoch (a strong tell that something is accumulating).

The fix — sever the graph by extracting a Python float:

total_loss += loss.item() # scalar → no graph reference
# or, if you need a tensor: loss.detach()

.item() (or .detach()) cuts the link so the graph can be garbage-collected after each backward. Same principle as not stashing raw output in a forward hook (pytorch-hooks).


Levers to reduce memory

  • torch.no_grad() for validation/inference — no graph built, activations freed immediately (see pytorch-training-loop).
  • Gradient checkpointing — recompute activations in backward instead of storing them; ~√N stored instead of N, ~33% extra compute (autograd-and-autodiff).
  • Lower precision (fp16/bf16) — halves parameter/gradient/activation bytes (see tensor-dtypes).
  • Smaller batch / gradient accumulation — trade per-step activation memory for more steps.
  • .item() / .detach() on anything you log or accumulate — never hold graph nodes you don’t need.

Related: autograd-and-autodiff, pytorch-training-loop, pytorch-hooks, tensor-dtypes, dataloader-and-batching