Autograd in Depth

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.

Read these first

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

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 away

The 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

Developer — Code and libraries.

Setup

bash
pip install torch

Outputs verified with torch 2.5.1, CPU.

Watching the recording happen

graph.py
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], "...")
Output
<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:

retain.py
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 replace
Output
tensor(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

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