Debugging PyTorch

Memory that grows every epoch

The classic PyTorch memory leak is storing a loss tensor instead of its number — each one drags its whole computation history along, and the drawer fills until the crash.

On this page 5
  1. Why it exists
  2. How it works
  3. Where you have seen this
  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.

When training uses more memory every epoch, something in your loop is keeping objects that should have been thrown away — usually without your knowledge.

Imagine keeping accounts by writing the day's total in a notebook. Sensible. Now imagine instead stapling the entire day's stack of receipts into the notebook each evening, because the total is technically in there somewhere. Within a month the notebook cannot close.

In PyTorch, a loss value comes attached to its receipts: the full record of every calculation that produced it. Store the loss itself, and you store the receipts. Store the plain number, and the receipts go in the bin where they belong.

Why it exists

To compute gradients, PyTorch keeps a computation graph — a record of every operation and its intermediate results — connecting your loss back to the weights. That record is bulky by design; it holds working values from every layer.

Normally, the record is discarded right after the backward pass. But records are only discarded when nothing points at them. Keep the loss in a list — for a plot, a log, a running total — and the loss keeps its record. That record keeps every intermediate value. One innocent line, gigabytes of stowaways.

How it works

history.append(loss)          ← keeps loss + its ENTIRE receipt trail
                                 memory: grows every step until the crash

history.append(loss.item())   ← keeps one plain number
                                 memory: flat forever

The difference is a single method: item converts a one-value tensor into an ordinary number, leaving the receipts behind.

Where you have seen this

Phones fill the same way: nobody decides to store four thousand photos of receipts and screenshots; each keep felt free at the moment. Growth-by-accumulation is always invisible per step and undeniable in total — which is why the diagnostic below counts, rather than trusts feelings.

Remember this

  • A loss tensor carries its whole computation record.
  • Store numbers (.item()), not tensors, when logging or plotting.
  • Memory that grows linearly with steps means the loop keeps something per step.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Captured with torch 2.5.1 on CPU; seeded and deterministic. The demo runs identically without a GPU — this leak is about the graph, not the device.

The leak, counted live

The counter walks Python's object registry and counts living tensors — crude, but wonderfully honest:

leaky_loop.py
import gc
import torch
import torch.nn as nn

def live_tensors():
    return sum(1 for obj in gc.get_objects() if type(obj) in (torch.Tensor, nn.Parameter))

torch.manual_seed(0)
X, y = torch.randn(64, 10), torch.randn(64, 1)
model = nn.Sequential(nn.Linear(10, 32), nn.ReLU(), nn.Linear(32, 1))
opt = torch.optim.SGD(model.parameters(), lr=0.01)
loss_fn = nn.MSELoss()

history = []
for epoch in range(1, 401):
    opt.zero_grad()
    loss = loss_fn(model(X), y)
    loss.backward()
    opt.step()
    history.append(loss)                 # the leak: a tensor drags its whole graph along
    if epoch % 100 == 0:
        print(f"epoch {epoch:3d}  live tensors: {live_tensors()}")
print("history[0] still carries a graph:", history[0].grad_fn is not None)
Output
epoch 100  live tensors: 110
epoch 200  live tensors: 210
epoch 300  live tensors: 310
epoch 400  live tensors: 410
history[0] still carries a graph: True

One new permanent tensor per epoch, in lockstep — plus the last line's confession: even epoch 1's loss still holds its grad_fn, the handle to its graph, four hundred epochs later.

Now change the append to history.append(loss.item()), and the final print to print("history[0] is now:", type(history[0]).__name__) — a float has no grad_fn to ask about. Rerun:

Output
epoch 100  live tensors: 11
epoch 200  live tensors: 11
epoch 300  live tensors: 11
epoch 400  live tensors: 11
history[0] is now: float

Eleven tensors: six model parameters, their six-minus-one grads and data — the working set — flat from epoch 100 to 400. That flatness is what a healthy loop looks like.

Where this bug hides in real code

total_loss += loss — the accumulator version. Each += builds a growing sum-graph chaining every batch of the epoch. Fix: total_loss += loss.item(), exactly as practised in the validation lesson.

Logging dicts — metrics["loss"].append(loss), tensors tucked into dashboards and printed later. Same disease, scattered locations.

Storing outputs for later analysis — all_preds.append(out) keeps each batch's graph during training. In an eval loop under torch.no_grad() there is no graph, so storing outputs there costs only the tensors themselves; storing out.detach().cpu() is the polite habit either way.

Keeping the batch across iterations — holding a reference to the previous batch (for contrastive tricks, or by accident in a closure) keeps its graph if that graph is still attached.

On a GPU, this same bug crashes faster and louder: CUDA out of memory mid-epoch, with the reported free memory shrinking each epoch. The mechanics of the caching allocator — why nvidia-smi disagrees with your arithmetic — are in moving tensors between CPU and GPU.

Common mistakes

Blaming the data or the model size. A model that fits in memory at epoch 1 fits at epoch 100 — model size produces constant usage. Linear growth is always an accumulation in the loop. The shape of the curve is the diagnosis.

"Fixing" it with gc.collect() or torch.cuda.empty_cache() every step. Neither frees objects your list still points at; both add overhead. They treat the symptom of a reference that should not exist.

Using detach() but keeping GPU tensors. loss.detach() drops the graph but the tensor still occupies GPU memory, and ten thousand of them add up. For logging scalars, .item(); for keeping arrays, .detach().cpu().

Calling .item() in tight per-step loops on GPU without need. The opposite over-correction: .item() forces a device synchronisation. Log every N steps, or accumulate on-device and read once per epoch — the balance discussed in that same tensors lesson.

Try it yourself

Plant the accumulator version (total_loss += loss, printing total_loss.item() per 100 epochs) and watch the counter climb again. Then inspect total_loss.grad_fn and follow total_loss.grad_fn.next_functions two levels down — you are looking directly at the chained receipts.

What to learn next

Researcher — Mathematics and papers.

Lifetime semantics

Autograd's graph is kept alive by ordinary Python/C++ reference counting: each output tensor owns its grad_fn node; nodes own their saved tensors (via SavedVariable) and edges to predecessor nodes. A single retained output therefore transitively pins the whole backward closure — the demo's history[0].grad_fn is not None is the observable root of that chain. After backward() (with default retain_graph=False) the engine releases saved tensors it consumed, but the node structure — and any saved tensors of ops not yet executed backward, e.g. in a partially-backwarded graph — persists while referenced. In-place ops complicate saved-tensor validity, which is the adjacent topic of in-place operations and autograd.

Activation memory arithmetic

Training memory is dominated by saved activations: $O!\left(B \sum_l d_l\right)$ for batch $B$ and layer widths $d_l$ — typically several times parameter memory for CNNs and transformers. One retained epoch-end graph therefore costs roughly one training step's activation footprint; an accumulated per-batch chain costs the epoch's. Gradient checkpointing (Chen et al., 2016, Training Deep Nets with Sublinear Memory Cost) trades this for recomputation, $O(\sqrt{L})$ retention — torch.utils.checkpoint — orthogonal to, and no cure for, reference leaks.

Instruments

  • CPU: the gc.get_objects census used here; tracemalloc for allocation-site attribution; objgraph.show_backrefs to render what holds the reference — the leak's paper trail as an actual graph.
  • CUDA: torch.cuda.memory_allocated() / max_memory_allocated() per epoch (linear growth = this lesson); torch.cuda.memory_summary() for allocator state; and since 2.x, torch.cuda.memory._record_memory_history() producing snapshot timelines renderable at pytorch.org/memory_viz — the definitive tool, attributing every live block to a Python stack.
  • Discipline: an epoch-boundary assertion — assert torch.cuda.memory_allocated() < budget or the tensor census — turns silent growth into a loud failure at the epoch it starts, which converts a whodunit into a diff of the last change.

The eval-loop corollary

A validation pass without torch.no_grad() builds graphs for every batch and — because no backward() ever runs to consume them — holds all of them until the loop's references drop. This is the single most common cause of "OOM only during validation", and the third guard of the validation lesson is its one-line prevention; torch.inference_mode() is the stricter variant.

What to learn next