Autograd Gotchas: detach, no_grad, retain_graph
Navigating common PyTorch autograd graph memory leaks, detach operations, no_grad contexts, and retain_graph errors.
PyTorch Dynamic Computational Graph Engine
PyTorch builds dynamic computational graphs on the fly during forward passes:
x ──► [ Linear Layer ] ──► z ──► [ ReLU Activation ] ──► a ──► [ Loss ]
(requires_grad=True) │
▼
loss.backward() (Frees Graph Buffers!)
Understanding how Autograd manages memory prevents GPU Out Of Memory (OOM) crashes and graph execution bugs.
3 Critical Autograd Gotchas
┌──────────────────────────┬──────────────────────────┬──────────────────────────┐
│ 1. MEMORY LEAK VIA LISTS │ 2. TENSOR DETACHING │ 3. RETAIN GRAPH ERRORS │
├──────────────────────────┼──────────────────────────┼──────────────────────────┤
│ Appending loss tensors to│ tensor.detach() breaks │ Calling backward() twice │
│ lists keeps full graph in│ graph history, creating a│ throws RuntimeError unless│
│ VRAM memory! │ static value tensor. │ retain_graph=True is set!│
└──────────────────────────┴──────────────────────────┴──────────────────────────┘
1. Memory Leaks from Appending Tensors (.item())
# BAD: Appending live tensor keeps full autograd graph history in VRAM!
losses = []
for inputs, targets in loader:
loss = criterion(model(inputs), targets)
losses.append(loss) # GPU OUT OF MEMORY CRASH!
# GOOD: Convert to scalar float using .item()!
losses = []
for inputs, targets in loader:
loss = criterion(model(inputs), targets)
losses.append(loss.item()) # Safely extracts scalar number, discarding graph!
2. Disconnecting Graphs (.detach())
tensor.detach() creates a new tensor that shares the underlying data memory but has requires_grad = False and no history in the computational graph.
Used in Generative Adversarial Networks (GANs):
- When updating the Discriminator on fake images generated by the Generator, detach fake images so Discriminator gradients do not update Generator weights!
fake_images = generator(noise)
# Detach fake images so backprop stops at Discriminator boundary
discriminator_loss = criterion(discriminator(fake_images.detach()), real_labels)
3. Multiple Backward Passes (retain_graph=True)
By default, PyTorch frees intermediate activation buffers immediately after loss.backward() finishes to save VRAM memory.
Calling backward() a second time on the same graph throws:
RuntimeError: Trying to backward through the graph a second time...
Fix by passing retain_graph=True:
loss1.backward(retain_graph=True) # Keeps graph buffers in memory
loss2.backward() # Frees graph buffers after execution
Disabling Autograd (torch.no_grad())
During inference and validation, wrap code inside with torch.no_grad()::
with torch.no_grad():
predictions = model(inputs) # Disables autograd graph construction!
Reduces GPU memory footprint by up to $50%$ and speeds up evaluation.
Say this out loud
Navigating PyTorch autograd requires managing graph references carefully. Appending raw loss tensors to lists causes memory leaks; call loss.item() to extract scalars safely. Calling tensor.detach() breaks graph history for adversarial or multi model updates. Using torch.no_grad() disables graph building during validation to save GPU memory.
Followups to expect
- What is
torch.inference_mode()vstorch.no_grad()?inference_mode()is a newer, faster context manager that disables autograd and view tracking entirely, offering even higher speedups thanno_grad(). - What is
requires_grad_()in-place method? Toggling tensor gradient tracking status in place, useful for freezing pre-trained model backbone layers during transfer learning (param.requires_grad = False).
Check yourself
Why does appending a tensor directly to a Python list inside a training loop cause severe GPU memory leaks?