In PyTorch, why does keeping the loss tensor grow memory each iteration?
answer
- grad_fn is a handle on a subgraph
- freed only when nothing references it
- metrics code, not model code
- monotonic growth, never recovers
- item or detach at the boundary
basics
~20 sA loss tensor holds a reference to the whole computation graph that produced it, including every saved activation. Storing it in a list or adding it to a running total keeps that graph alive, so each iteration's activations pile up. Use loss.item() or loss.detach() instead.
solid answer
~50 sAfter a backward pass PyTorch frees the buffers the graph saved — but only if nothing still references the graph. A loss tensor *is* such a reference: through `grad_fn` it points at the entire chain of backward nodes, and those nodes hold the activations they need for their derivatives. So `total += loss` keeps iteration 1's activations alive at iteration 200, and memory climbs linearly until you hit an out-of-memory error, often after the run has looked healthy for a while. The fix is to strip the graph off anything you retain: `total += loss.item()` for a Python float, or `.detach()` when you need the tensor. The same failure appears with a list of per-batch losses, a logging buffer, an RNN hidden state carried across steps without detaching, and code that passes `retain_graph=True` to silence the "backward through the graph a second time" error rather than fixing the cause.
code
python · 11 linesimport torch
model = torch.nn.Linear(4, 1)
running = 0.0
for _ in range(3):
loss = model(torch.randn(8, 4)).pow(2).mean()
loss.backward()
running += loss.item() # float: this step's graph can be freed
# running += loss # keeps every step's graph alive
print(loss.grad_fn is not None) # True: the tensor holds a subgraph
print(running)go deeper
Know the rule of thumb: use loss.item() when logging or summing losses, because keeping the tensor itself keeps the whole computation graph alive.
Explain that a tensor's grad_fn references saved activations, that backward frees them only when nothing else holds the graph, and why a list of losses or an undetached hidden state leaks.
Diagnose it in production: monotonic per-iteration growth in torch.cuda.memory_allocated(), audit every object outliving the loop for a non-None grad_fn, and refuse retain_graph=True as a fix until you know why the graph is traversed twice.
Make it structural rather than tribal knowledge: a shared training loop that detaches at the metrics boundary, a memory-per-step assertion in CI, and review guidance that treats retain_graph=True as requiring a written justification.
## What a loss tensor really holds A scalar loss looks like one number, but it carries `grad_fn`, which points at a backward node, which points at *its* inputs' backward nodes, and so on all the way to the parameters. Those nodes hold saved tensors: the activations, intermediate results and sometimes the inputs they need to compute their local derivatives. For a realistic network that is by far the largest allocation of the step — usually much bigger than the parameters themselves. PyTorch frees those saved buffers at the end of `backward()`, because in the normal case nobody needs them again. But freeing is reference-counted, not scheduled: if a Python object still points at the graph, nothing is released. Holding the loss holds the graph holds the activations. ## The canonical bug ``` running = 0.0 for batch in loader: loss = criterion(model(batch.x), batch.y) loss.backward() optimizer.step(); optimizer.zero_grad() running += loss # <-- the leak ``` `running` starts as a float, but `float + tensor` produces a *tensor with a grad_fn*, and now every iteration's graph is chained into one ever-growing structure. Memory rises linearly with iteration count. On GPU you get `CUDA out of memory` somewhere in the middle of an epoch — with a message pointing at whichever allocation happened to be unlucky, not at the line that caused it. The two correct spellings: - `running += loss.item()` — pulls the scalar out as a Python float. Note it forces a device synchronisation, so in a very tight loop prefer accumulating on device with `.detach()` and calling `.item()` once at the end. - `losses.append(loss.detach())` — keeps a tensor but severs the graph. The same rule applies to anything you stash: predictions kept for a validation metric, attention maps saved for visualisation, per-example losses collected for analysis. If it came out of a tracked forward pass and you are keeping it beyond the iteration, detach it. ## The other shapes this takes **Recurrent hidden state.** Carrying `h` from one chunk to the next without `h = h.detach()` extends the graph across chunk boundaries indefinitely — the graph grows until backward becomes both slow and enormous. Truncated backpropagation through time is precisely the discipline of detaching at chunk boundaries. **retain_graph=True used as a painkiller.** Calling backward twice on one graph raises `RuntimeError: Trying to backward through the graph a second time, but the saved intermediate results have already been freed`. The advice found in search results is to pass `retain_graph=True`, and it does make the error go away — while keeping every graph alive. Sometimes retention is genuinely needed (two losses over one shared forward pass, some GAN formulations), but far more often the real bug is a graph accidentally carried between iterations. Ask *why* the same graph is being traversed twice before reaching for the flag. **create_graph=True.** Needed for second-order gradients, but it makes `.grad` tensors themselves carry graphs. Clear or detach them, or the accumulator becomes the leak. **Hooks and closures.** A backward hook or a callback that captures the loss (or any intermediate) in an enclosing scope pins the graph just as effectively as a list does. ## Diagnosing it The signature is distinctive: memory grows *monotonically per iteration* rather than spiking at a particular batch. A batch-size or sequence-length spike causes a sawtooth that recovers; a retained graph never recovers. Steps that work: 1. Print `torch.cuda.memory_allocated()` at the same point in every iteration. A straight upward line means retention. 2. Compare with `torch.cuda.max_memory_allocated()` to separate a genuine leak from fragmentation, and `torch.cuda.memory_summary()` for a breakdown. 3. Grep the loop for anything that survives it — lists, dicts, `self.` attributes, running totals — and check whether each holds a tensor with a `grad_fn`. `t.grad_fn is not None` on a stored object is the smoking gun. 4. Bisect by replacing the suspect accumulation with `.item()` and re-running for a few hundred steps. ## The mental model to carry Say it as a rule: *a tensor with a `grad_fn` is a handle on an entire subgraph, not a number.* Everything else follows — why `.item()` and `.detach()` are the standard hygiene at the boundary between training and logging, why the leak shows up in metrics code rather than model code, and why the error surfaces far from its cause.
- What does the error about backward through the graph a second time actually mean?After a backward pass, autograd frees the intermediate results its nodes saved, so traversing the same graph again has nothing to work with. The error means you called backward twice over one forward pass. Legitimate cases exist — two losses over a shared forward — and there retain_graph=True is correct. More often it signals that a tensor from a previous iteration is still wired into this one, and retaining the graph converts an error into a memory leak.
- How does detaching an RNN hidden state between chunks relate to this?Carrying the hidden state forward without detaching links each chunk's graph to the previous one, so the graph — and its saved activations — grows without bound and backward gets progressively slower. Calling h = h.detach() at the chunk boundary keeps the numeric state while cutting the history, which is exactly what truncated backpropagation through time means. The model still sees the carried state; gradients simply stop at the boundary.
- When is retain_graph=True genuinely the right answer?When you deliberately need more than one backward pass over the same forward computation: two loss terms computed from shared intermediates that you cannot or do not want to sum before back-propagating, some adversarial setups, and gradient-penalty formulations that differentiate through an already-differentiated quantity. Use it knowingly and scope it tightly, since every retained graph is held activation memory. If the second traversal is accidental, retaining it hides the bug.
- How do you tell a retained-graph leak from ordinary fragmentation?Sample torch.cuda.memory_allocated() at the same point each iteration. A retained graph produces a monotonic climb that never returns to baseline. Fragmentation shows as allocated memory staying flat while reserved memory grows or allocation fails despite apparent headroom — compare memory_allocated() with memory_reserved(). Variable sequence lengths give a sawtooth that recovers, which is a third pattern and is addressed by bucketing rather than by detaching.
saying these in an interview costs you the question
- Adding a loss tensor to a running total instead of loss.item()
- Passing retain_graph=True to make an error disappear
- Blaming the batch size when memory grows monotonically every iteration
- Thinking backward() always frees the graph regardless of references
- Carrying an RNN hidden state across chunks without detaching