How the autograd graph is built and freed
Every calculation on a tracked tensor is recorded into a graph, backward() walks that recording in reverse to produce gradients, and then the recording is destroyed.
- 7 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
While your code computes forwards, PyTorch secretly writes down every step — and later reads the notes backwards to work out what to blame.
Think of trekking into a forest and tying a ribbon to a branch at every turn. Walking in, the ribbons cost you almost nothing. Walking out, you follow them in reverse and they lead you exactly home. Then you collect the ribbons — they were for one trip only.
That trail of ribbons is the autograd graph: a recording of every calculation, kept so it can be replayed backwards.
Why it exists
Training needs to answer: "how much did each weight contribute to the error?" — for millions of weights. Working that out by hand, for whatever model you happen to write, is impossible to maintain.
So PyTorch watches you instead. Whatever you compute — any layers, any maths, any if-statements — it records what happened. When you call for blame to be assigned, it follows its own notes backwards. You never write the backwards logic; the recording is the logic. This is why the graph is called dynamic: it is rebuilt fresh from whatever your code actually did, every single run.
How it works
forward: x ──(square)──> y ──(add 1)──> z each arrow is recorded
backward: x <──(blame)─── y <──(blame)── z notes read in reverse
then: the notes are thrown awayThe throwing away matters. The trail serves one backward walk, and keeping trails around would eat memory fast. Next forward run, a fresh trail is tied.
A real example you have seen
Google Maps does the forward-and-back trick on every trip. It tracks the route you took; ask for the way back and it replays the route in reverse. It does not store every road in the country for you — one recording per journey, then it moves on.
Remember this
- Computing forwards quietly builds a recording of every step.
backward()reads the recording in reverse to assign blame (gradients).- The recording is then freed — one recording, one backward walk.
What to learn next
- requires_grad and leaf tensors — who gets a
.gradat the end of the walk, and why. - Backpropagation — the algorithm this machinery automates.
- detach, no_grad and inference_mode — switching the recorder off on purpose.
Developer — Code and libraries.
Setup
pip install torchOutputs verified with torch 2.5.1, CPU.
Watching the recording happen
import torch
x = torch.tensor(3.0, requires_grad=True)
y = x * x
z = y + 1.0
print(z.grad_fn) # the last recorded step
print(z.grad_fn.next_functions) # what it connects back to
z.backward() # walk the recording backwards
print(x.grad) # dz/dx = 2x = 6
try:
z.backward() # the recording was thrown away
except RuntimeError as err:
print("second backward:", str(err)[:80], "...")<AddBackward0 object at 0x0000015276A01D50> ((<MulBackward0 object at 0x0000015276A000D0>, 0), (None, 0)) tensor(6.) second backward: Trying to backward through the graph a second time (or directly access saved ten ...
The hex addresses will differ on your machine; the structure will not.
The walkthrough
grad_fn is the recording made visible. z was made by an addition, so its grad_fn is AddBackward0. Follow next_functions and you find MulBackward0 — the squaring — and a None for the constant 1.0, which needs no blame. Your whole computation is findable this way, link by link. A tensor you created has grad_fn of None: recordings start at operations, not at inputs.
requires_grad=True is the on-switch. Only computations touching a tracked tensor get recorded. Without it, y = x * x produces a plain tensor with no history — try it, y.grad_fn prints None.
The second backward() fails on purpose. Walking the recording consumes it — the saved values needed for the walk are freed immediately to reclaim memory. This error is nearly always a design smell rather than a missing flag, but when you genuinely need two walks, say so:
import torch
x = torch.tensor(3.0, requires_grad=True)
z = x * x
z.backward(retain_graph=True) # keep the recording alive
z.backward() # allowed now
print(x.grad) # 6 + 6: gradients ADD, they do not replacetensor(12.)
And note what it printed: 12, not 6. Gradients accumulate into .grad. Every training loop you have seen calls optimizer.zero_grad() for exactly this reason — without it, every step's blame piles onto the last step's. Accumulating on purpose is a real technique covered in gradient accumulation.
Common mistakes
Meeting "Trying to backward through the graph a second time" in a loop. Usually some tensor from a previous iteration — often a running loss like total = total + loss — still links into the old, freed recording. Accumulate numbers, not tracked tensors: total += loss.item(), or .detach() it.
Forgetting zero_grad(). The loss wanders instead of falling, because gradients from every step are summing. This is the first thing to check when a fresh training loop misbehaves.
Recording during evaluation. Computing validation loss without switching recording off builds trails you will never walk, at real memory cost. The off-switches — no_grad and friends — get their own lesson: detach, no_grad and inference_mode.
Keeping a tracked tensor alive keeps its whole recording alive. Storing loss in a Python list stores the entire graph behind it, every iteration. Memory climbs each epoch; the fix is again .item() or .detach().
Try it yourself
Build z = (x * y) + (x * 2) with two tracked scalars, print z.grad_fn.next_functions, and draw the graph on paper before running backward(). Predict both .grad values, then check. Then add an if z > 5 branch that changes the formula, run with different inputs, and confirm the recording differs per run — that is the dynamic graph earning its name.
What to learn next
- requires_grad and leaf tensors — who gets a
.gradat the end of the walk, and why. - Backpropagation — the algorithm this machinery automates.
- detach, no_grad and inference_mode — switching the recorder off on purpose.
Researcher — Mathematics and papers.
Reverse-mode automatic differentiation
Autograd implements reverse-mode AD over a dynamically-constructed DAG. Forward execution of primitives f_1 … f_k builds nodes storing (a) the backward function of each primitive and (b) whatever operands or results that backward needs ("saved tensors"). Calling backward seeds the output adjoint and propagates:
x̄ = Σ_{p ∈ consumers(x)} (∂f_p/∂x)ᵀ ȳ_p
Where x̄ denotes the adjoint (gradient of the scalar loss with respect to x), ȳ_p the adjoint of consumer p's output, and ∂f_p/∂x the local Jacobian. Each node computes a vector–Jacobian product, never a full Jacobian — the fan-in sum over consumers is the accumulation behaviour observed above, and it is the multivariate chain rule, not a convenience choice.
Reverse mode computes gradients of one scalar with respect to n inputs in one forward plus one backward pass — O(1) passes, roughly a constant multiple of forward cost (the "cheap gradient principle", Griewank and Walther, 2008, Evaluating Derivatives). Forward-mode's O(n) passes is why reverse mode owns deep learning; the comparison becomes practical in Jacobians, vjp and jvp.
Define-by-run versus define-then-run
PyTorch's tape is rebuilt per iteration ("define-by-run"), inherited from Chainer (Tokui et al., 2015) and formalised for PyTorch in Paszke et al. (2017, Automatic differentiation in PyTorch, NIPS-W) and Paszke et al. (2019, NeurIPS). Control flow needs no special handling because only the taken branch exists in the tape. The cost is per-iteration graph construction overhead and no whole-graph optimisation — the gap torch.compile closes by tracing the tape into a static graph when shapes and branches permit.
Memory: the real constraint
Saved-for-backward activations dominate training memory: O(Σ intermediate sizes), typically dwarfing parameters. Eager freeing during the backward walk (the source of the "second time" error) is one mitigation; gradient checkpointing (Chen et al., 2016, Training Deep Nets with Sublinear Memory Cost) is the systematic one — save only O(√k) of k layers' activations and recompute the rest during backward, trading ~33% extra compute for order-of-magnitude activation memory savings.
retain_graph=True disables the freeing globally for that graph; torch.autograd.grad (functional form, no .grad mutation) plus explicit graph structure is usually the better tool where multiple backward passes are genuinely needed (GANs, some meta-learning).
References
- Griewank and Walther (2008), Evaluating Derivatives, 2nd ed., SIAM — the AD reference text.
- Baydin et al. (2018), Automatic Differentiation in Machine Learning: a Survey, JMLR 18 — the field map.
- Paszke et al. (2019), PyTorch: An Imperative Style, High-Performance Deep Learning Library, NeurIPS.
- Chen et al. (2016), arXiv:1604.06174 — checkpointing.
What to learn next
- requires_grad and leaf tensors — who gets a
.gradat the end of the walk, and why. - Backpropagation — the algorithm this machinery automates.
- detach, no_grad and inference_mode — switching the recorder off on purpose.