In the previous lecture we derived backpropagation by hand for a specific network. Nobody does that for a 100-layer transformer. Instead, frameworks like PyTorch, JAX and TensorFlow automatically differentiate any program built from differentiable operations. Understanding how they do it will make you a far better debugger and will demystify errors like "one of the variables needed for gradient computation has been modified by an inplace operation".
Three ways to compute derivatives#
- Numerical differentiation — finite differences $\frac{f(x + h) - f(x - h)}{2h}$. Easy but approximate and costs $O(P)$ evaluations for $P$ parameters. Useful only for checking.
- Symbolic differentiation — manipulate formulas like a computer-algebra system. Exact, but expressions can blow up in size ("expression swell") and it struggles with loops and branches.
- Automatic differentiation (autodiff) — decompose the program into elementary operations with known derivatives and apply the chain rule numerically as the program runs. Exact to floating-point precision and efficient.
Computational graphs#
Any computation can be written as a directed acyclic graph whose nodes are elementary operations. For
the graph is: $a = x + y$, $b = \sin(x)$, $f = a\cdot b$. Each node knows its local derivative with respect to its inputs: $\partial a/\partial x = 1$, $\partial b/\partial x = \cos x$, $\partial f/\partial a = b$, $\partial f/\partial b = a$.
Forward mode#
Forward mode propagates derivatives with the computation, from inputs to outputs, carrying a "tangent" $\dot{v} = \partial v/\partial x$ for one chosen input direction. It computes a Jacobian–vector product $\mathbf{J}\mathbf{v}$ in one pass. It is efficient when there are few inputs and many outputs. It can be implemented elegantly with dual numbers $a + b\varepsilon$ where $\varepsilon^2 = 0$.
Reverse mode#
Reverse mode first runs the program forward, recording the graph and intermediate values. Then it propagates adjoints $\bar{v} = \partial L/\partial v$ backwards from the output:
It computes a vector–Jacobian product $\mathbf{u}^\top\mathbf{J}$ — the gradient of one scalar with respect to all inputs — in a single backward pass. Since training has one scalar loss and millions of parameters, reverse mode is exactly what we need. Backpropagation is reverse-mode autodiff applied to neural networks.
| Forward mode | Reverse mode | |
|---|---|---|
| Computes | $\mathbf{J}\mathbf{v}$ (JVP) | $\mathbf{u}^\top\mathbf{J}$ (VJP) |
| Cost per pass | ~1 forward | ~1 forward + 1 backward |
| Efficient when | inputs ≪ outputs | outputs ≪ inputs (e.g. a loss) |
| Memory | Low | Stores intermediate values |
Build a tiny autodiff engine#
In the spirit of Andrej Karpathy's micrograd:
import math
class Value:
def __init__(self, data, parents=(), op=""):
self.data, self.grad = data, 0.0
self._parents, self._op = parents, op
self._backward = lambda: None
def __add__(self, other):
other = other if isinstance(other, Value) else Value(other)
out = Value(self.data + other.data, (self, other), "+")
def _backward():
self.grad += out.grad
other.grad += out.grad
out._backward = _backward
return out
def __mul__(self, other):
other = other if isinstance(other, Value) else Value(other)
out = Value(self.data * other.data, (self, other), "*")
def _backward():
self.grad += other.data * out.grad
other.grad += self.data * out.grad
out._backward = _backward
return out
def sin(self):
out = Value(math.sin(self.data), (self,), "sin")
def _backward():
self.grad += math.cos(self.data) * out.grad
out._backward = _backward
return out
def backward(self):
order, seen = [], set()
def topo(v):
if v not in seen:
seen.add(v)
for p in v._parents:
topo(p)
order.append(v)
topo(self)
self.grad = 1.0
for v in reversed(order): # reverse topological order
v._backward()
x, y = Value(2.0), Value(3.0)
f = (x + y) * x.sin()
f.backward()
print("f =", round(f.data, 4))
print("df/dx =", round(x.grad, 4), " expected:", round(math.sin(2) + (2 + 3) * math.cos(2), 4))
print("df/dy =", round(y.grad, 4), " expected:", round(math.sin(2), 4))Note the += in each backward function: when a variable is used in several places ($x$ appears twice here), gradients from all paths accumulate — the "sum over paths" in the multivariable chain rule.
How PyTorch does it#
PyTorch builds a dynamic graph ("define-by-run"): each tensor operation records a node with a grad_fn as your Python code executes, so ordinary control flow (loops, if-statements) just works. Calling .backward() traverses the recorded graph in reverse.
import torch
x = torch.tensor(2.0, requires_grad=True)
y = torch.tensor(3.0, requires_grad=True)
f = (x + y) * torch.sin(x)
print(f.grad_fn) # <MulBackward0 ...>
f.backward()
print(x.grad, y.grad)Practical consequences:
- Gradients accumulate in
.grad— calloptimizer.zero_grad()every step. torch.no_grad()disables graph recording for inference, saving memory..detach()cuts a tensor out of the graph (used for target networks, stop-gradient tricks).- In-place operations can overwrite values needed for the backward pass, hence the famous error.
- The graph is freed after
backward()unlessretain_graph=True.
JAX takes a functional approach: jax.grad(f) returns a new function computing the gradient, composable with jax.jit (compilation), jax.vmap (vectorisation) and jax.jvp/jax.vjp for forward and reverse mode.
Memory tricks#
Reverse mode stores activations, which dominates memory for large models. Gradient (activation) checkpointing stores only some activations and recomputes the rest during the backward pass — trading about one extra forward pass for a large memory reduction, which lets much bigger models or batches fit on a GPU.