Automatic Differentiation Primer

loss.backward() is one line, and it returns a gradient for every one of a billion parameters at the price of about two forward passes. This page takes it apart — the tape, the two sweep directions, the memory, and the three places the answer is not the derivative you meant. Every number is computed by the figure beside it.

01

Three ways to get a slope

A model is a program with a billion knobs, and training needs to know which way each one moves the loss.

There are three ways to get a derivative out of a function you have written down: differentiate it by hand, nudge the input and watch the output, or differentiate the program. The third is autograd. The first two are why it exists.

First the thing itself. A derivative is a slope: put a point on the curve and the tangent through it carries one number — how fast f climbs there. Drag the point and watch that number follow:

x = 2.00 · f′ = 28. drag the point along the curve; the arrow keys move it 0.1 at a time and Home returns it to x = 2
x = 2.00 · f′ = 28

Notice that the tangent is the one line that touches the curve and lies along it. At x = 2 it reads 28 while f reads 49 — two facts about one point, and training wants the slope.

Getting that number is the problem. The definition offers a route: take a second point h to the right, draw the line through both, and shrink h. Close the gap and watch the two lines fall together:

h = 1 · 32

At h = 1 the secant reads 32 against the tangent's 28, and each notch narrows the gap — here the error is exactly 4h. So take h as small as the machine holds and the answer comes out exact. It does not.

Every value below is real float64, not a model of it. Shrink h past and the error stops falling and climbs again. Both axes are logarithmic — one tick is a factor of ten:

h = 1 · 0 digits
h = 1 · 0 digits

Watch where the curve turns. Its left arm falls because the secant is still a chord; its right arm climbs because f(x+h) and f(x) share more and more leading digits, and subtracting them throws those digits away. The best step is 1e−8 — near √ε — and it buys 8 correct digits of the 16 we started with. At , x + h is x and the answer is zero.

Losing half the digits would be survivable. The cost is not. Take the plainest function of n inputs, f = x₁·x₂·…·xₙ, and count the multiplications each method needs for the whole gradient — by hand, by finite differences, or swept in reverse:

n = 4 · reverse is still behind
n = 4 · reverse is still behind

The hand-written gradient rebuilds a product per component, so it costs n(n−2). Finite differences needs n+1 evaluations — about the same, and wrong in the eighth digit besides. The reverse sweep costs 3n−2, and is genuinely behind until n = 5. At it is 341× ahead, and a language model has n in the billions.

02

A product along a path

One rule does all of the work, and everybody already learned it.

Differentiating a program sounds harder than differentiating a formula. It is easier: a program is already broken into pieces small enough that somebody has written down the derivative of each piece.

Here is y = (2x+3)² as three operations. The value rides each wire left to right; under every box sits that operation's own derivative, evaluated where the value landed. Drag x and watch both rows move:

x = 2 · dy/dx = 28

Notice that the three local derivatives are the only calculus on this page: 2 for a doubling, 1 for adding a constant, 2v for a square. None of them knows the others exist. The chain rule says the answer is their product — 2 · 1 · 14 = 28 — and that product is a path from x to y.

A product multiplies in either order, and that one fact is the whole of what follows. Carry a running number from the left: start at 1 and multiply in each local derivative as you cross it. Press play, or walk it one operation at a time:

after 0 of 3 ops · 1

That is forward mode, and the number it carries is ẋ — how fast this wire moves when x moves. It travels beside the values, in the same direction, so nothing has to be stored.

Now the same three multiplications, right to left. Start at the output with 1 and push backwards; the number carried is x̄ — how much y moves when this wire moves. Play it, or :

after 0 of 3 ops · 1

Read the two sweeps together: 1, 14, 14, 28 backwards against 1, 2, 2, 28 forwards. Same three factors, same answer, opposite association. For one input and one output the two modes are exactly as good as each other.

03

The tape

Going backwards means the forward pass has to leave something behind.

Forward mode never looks back. The reverse sweep starts at the end, so every local derivative it wants was computed on the way there — and something has to have remembered them.

So the framework writes a row as each operation runs: which op, what went in, what came out, and what its backward rule will need kept. Run the program and watch the tape grow — it starts empty:

tape empty

Notice the last column. +3's backward rule is "pass the gradient through", which needs nothing from the forward pass; (·)²'s rule is 2v, so it keeps its one input. That column is the entire memory cost of autograd.

All of it is twelve lines of Python — a list, an operation that appends to it as it computes, and a loop that walks the list backwards:

tape = []            # (out, ins, locals)

def mul(a, b):
    out = V(a.v * b.v)
    tape.append((out, (a, b), (b.v, a.v)))
    return out

def backward(y):
    y.g = 1.0
    for out, ins, loc in reversed(tape):
        for x, d in zip(ins, loc):
            x.g += out.g * d   # += , never =

With the tape written, backward is that for loop over it in reverse. Each row multiplies the adjoint arriving at its output by its own local derivative and pushes the result down to its inputs. Read the tape from the bottom:

seed ȳ = 1

The invariant, and it is the whole correctness argument: the rows are replayed in reverse execution order, so every consumer of a value is finished before the value is read — which means its adjoint already holds ∂y/∂v summed over every path from v to y. As a loop assertion: v.grad == sum(c.grad * dc_dv for c in consumers(v)).

That last column has a price, and not the same one per operation. Measured in one 8 × 1024 × 768 fp16 activation — GPT-2's shape, 12.00 MiB — what a rule must keep spans more than an order of magnitude. Switch between four:

add · 0 B

add is free: its rule is a constant. relu needs only the sign of its input — one bit per element, 0.75 MiB. square keeps its input at 12.00 MiB, and keeps both operands at 24.00 MiB. Trading a multiply for an add trades part of this table.

One thing the straight chain hid. Let x be used twice — y = x²·(x+1) — and there are two paths from x to y. Drag x, watch the values go out and what each path carries back come in, and see what the node they meet at does with them:

x = 2 · x̄ = 16

The node adds. At x = 2 the paths carry 12 and 4 and x̄ comes out 16, which is 3x² + 2x. That += is why a sum over paths comes out right with nobody enumerating paths — and in §06 it is why one missing line ruins a training run.

04

One sweep, one question

A sweep is cheap. What it is not is general.

Everything so far had one input and one output, the one shape where the two modes tie. Take a program with two inputs and three outputs instead: its derivative is not a number but a 3 × 2 block, one partial per output-and-input pair.

Forward mode carries one seed at a time. Set ẋ₁ = 1, ẋ₂ = 0 and one sweep tells you how all three outputs move when x₁ moves — and nothing about x₂. Switch the seed and watch which wires light up:

ẋ₁ = 1 · one column

Each seed produces three numbers, which is one column of the block. So this 3 × 2 Jacobian takes two forward sweeps, one per input, and there is no clever seed that gets both at once: the sweep is linear in the seed, and two independent columns need two independent seeds.

Scaled up, that is the whole cost model. Here is a 4 × 6 Jacobian — six inputs, four outputs — with one forward sweep per column. Push the slider and fill it:

0 of 6 sweeps

Six sweeps for six columns: the count is the number of inputs, and nothing about the outputs enters it. A program with four outputs and one with four hundred cost forward mode the same.

§02's reverse sweep is that picture turned ninety degrees. Seed one output with 1, push backwards, and what you learn is how that output responds to every input at once — a row, not a column:

0 of 4 sweeps

Four sweeps instead of six, because the count is now the number of outputs. Nothing was optimised: same chain rule, same local derivatives, same multiplications per sweep. Only the association order changed, and with it which dimension you pay for.

05

Why the reverse sweep wins

Training is a billion numbers in and one number out.

A loss is a scalar, so loss.backward() is a request for a 1 × n Jacobian — one row, n columns, n the parameter count. §04 already priced both ways of filling it.

Both blocks below are that Jacobian, filled column by column on the left and row by row on the right. Set the number of inputs and of outputs, and read off which finishes first:

n = 8 · m = 1

At 8 inputs and 1 output the reverse block finishes in one sweep where the forward block needs eight. Flip the shape — against many outputs — and forward wins by the same argument. Sweep along the short side.

For a loss the short side is never in doubt. Plot the sweeps each mode needs against the number of inputs, holding outputs at one. Both axes are logarithmic, so a straight diagonal is a straight proportionality:

n = 1
n = 1

Reverse is the flat line at 1. Forward is the diagonal. At that is one sweep against a million, and a 7-billion-parameter model would need 7 billion forward sweeps for what one reverse sweep gives.

A sweep is not free, but it is bounded. Griewank's cheap-gradient result caps a reverse gradient at four function evaluations however many inputs there are; in a Transformer the measured split is one unit forward and about two back. Move the parameter count and watch the bars refuse to move:

125M · 3×

A training step is 3× a forward pass and provably under 4×, at 125M parameters and at 405B alike. That is where 6·N·P FLOPs per training token against 2·N·P for inference comes from — arithmetic, not coincidence. The gradient does not cost more because there is more to differentiate.

It costs something else. The tape must hold every saved activation from the whole forward pass before the first backward row can be read, so memory grows with depth — unless you throw some away and recompute it. On a 96-layer stack, keep every k-th boundary:

k = 1
k = 1 · 97 / 96

Without checkpointing the peak is 96 layers. Checkpointing at brings it to 20 — 4.8× less — because you store ⌈96/k⌉ boundaries and hold k layers live while replaying a segment, and that sum is smallest near √96. The price is one extra forward pass, a third more compute per step. So the level-two answer to "reverse mode is one sweep" is: one sweep of time, and a whole tape of space.

06

Exact, and about what

Autodiff is exact to the last bit. It is worth being precise about what it is exact about.

Every number so far came out with no truncation and no cancellation. But what a tape differentiates is the program that ran, and the program that ran and the function you meant are not always one object.

Start with the commonest activation in deep learning. relu has no tangent at zero — slope 0 on one side, 1 on the other, and no single line at the corner. Drag the point onto the corner and read what the framework hands back:

x = 1.40 · relu′ = 1. drag the point along the curve; the arrow keys move it 0.1 at a time and Home returns it to x = 1.4
x = 1.40 · relu′ = 1

It returns 0, by convention, without warning you. Any of 0, 1 or ½ is a defensible subgradient; PyTorch picks 0. It matters approximately never — landing exactly on 0 in floating point is negligibly likely — and enormously the day someone builds a rule out of x - x and wonders where the gradient went.

The serious version is control flow. A Python if puts the branch that ran on the tape and nothing else, so the test itself never becomes a row and has no derivative. Drag the point across the jump:

x = 1.40 · y = 1 · dy/dx = 0. drag the point across the jump; the arrow keys move it 0.1 at a time and Home returns it to x = 1.4
x = 1.40 · y = 1 · dy/dx = 0

Watch the readout refuse to move. y goes from 0 to 1 as the point crosses and dy/dx reads 0 on both sides, because the comparison was never recorded. This fails silently: the model trains, the loss falls for other reasons, and the discrete decision was invisible to every gradient you computed.

The fix is to differentiate a different function on purpose. Replace the step with σ(x/τ), which has a gradient everywhere, and pay a temperature for it. Lower τ and watch the gradient:

τ = 1

As τ falls, the curve approaches the step it stands in for — and its gradient collapses into a spike of width about τ, so narrow that at almost every sample lands outside it and learns nothing. Straight-through estimators and Gumbel-softmax manage that trade; neither removes it.

The last one is a footgun, and it is §03's +=. A gradient buffer accumulates by design, so it accumulates whether or not you meant it. Run four batches without the call, then switch it on:

after batch 0 · 0×

Both loops run clean. Without zero_grad(), .grad holds the sum of every batch so far, so batch 4 takes a step four times too large — and it does not crash, it trains badly, which is the expensive kind of bug. Two of this page's three failures are silent like that; mutating a saved tensor in place is the third, and it fails loudly, because the tape's version counter catches it.

07

The whole run

Both halves of the technique, end to end, under one control.

§03's twelve lines, running on §03's diamond: three operations forward, each appending a row, then four edge replays backward reading them off. Play it, or walk it a step at a time:

nothing has run

dy/dx = 16, out of a tape three rows long. The per-op rule was the two numbers in locals; there is no calculus anywhere else on this page, no step size, and no cancellation.

Everything a real framework adds is engineering on top of exactly that: rules in C++ instead of Python, tensors instead of floats, a graph instead of a list once the program branches, and §05's checkpointing so that the tape still fits in memory.