7.5 backward(): let the graph do calculus
microgpt has 4,192 knobs. To learn, it must know, for every single knob, “if I turn you a tiny bit, does the loss go up or down, and by how much?” Nudging each knob one at a time would mean 4,192 extra runs per step. Today you write the 14 lines that get all 4,192 answers in one backwards sweep.
1Watch
During the forward pass every Value remembers two things: its children (the values it was made from) and its local gradients, the “exchange rates” from each child to itself. For e = a × b, nudging a by 1 moves e by b, so the local gradient for a is b.data.
backward() does three things:
- Sort the graph so every node comes after its children. That’s the topological sort you wrote in lesson 3.3 (socks before shoes).
- Set the loss’s own gradient to 1.
- Walk the sorted list backwards. Each node passes its gradient to each child:
child.grad += local_grad × node.grad. That’s the chain rule: multiply the exchange rates along the path.
Two details carry the whole idea. Backwards order guarantees a node has received gradient from everyone who uses it before it passes anything on. And +=, not =: when a value is used in two places, both paths push on the loss, so their effects add up. In the widget, press “try = instead of +=” and step through: e ends at −6 instead of −2, and a and b inherit the mistake.
2Explore
Forward pass done: every node knows its value. Every grad starts at 0.
3Build
The Value class below is copied from microgpt, with two lines of backward() removed:
- In
build_topo: addvtotopoafter all of its children have been visited. That means after theforloop over children, at the same indentation as thatfor: not before it, and not inside it. - In the backwards loop: grow the child’s gradient by
local_grad * v.grad. Think hard about=vs+=.
The last check compares your gradients against plain nudging (the numerical derivative) on a messy formula with exp, log and relu. If they agree, your backward() is correct, full stop.
Stuck? In the code-along video, backward() starts at 2:30 and the step-by-step trace at 3:43. The earlier part rebuilds the Value class from lesson 7.4.
Stuck? Watch the code being written, line by line (6 min)
import math
class Value:
__slots__ = ('data', 'grad', '_children', '_local_grads') # Python optimization for memory usage
def __init__(self, data, children=(), local_grads=()):
self.data = data # scalar value of this node calculated during forward pass
self.grad = 0 # derivative of the loss w.r.t. this node, calculated in backward pass
self._children = children # children of this node in the computation graph
self._local_grads = local_grads # local derivative of this node w.r.t. its children
def __add__(self, other):
other = other if isinstance(other, Value) else Value(other)
return Value(self.data + other.data, (self, other), (1, 1))
def __mul__(self, other):
other = other if isinstance(other, Value) else Value(other)
return Value(self.data * other.data, (self, other), (other.data, self.data))
def __pow__(self, other): return Value(self.data**other, (self,), (other * self.data**(other-1),))
def log(self): return Value(math.log(self.data), (self,), (1/self.data,))
def exp(self): return Value(math.exp(self.data), (self,), (math.exp(self.data),))
def relu(self): return Value(max(0, self.data), (self,), (float(self.data > 0),))
def __neg__(self): return self * -1
def __radd__(self, other): return self + other
def __sub__(self, other): return self + (-other)
def __rsub__(self, other): return other + (-self)
def __rmul__(self, other): return self * other
def __truediv__(self, other): return self * other**-1
def __rtruediv__(self, other): return other * self**-1
def backward(self):
# Step 1: put every node in an order where a node comes AFTER
# all of its children (topological order - socks before shoes).
topo = []
visited = set()
def build_topo(v):
if v not in visited:
visited.add(v)
for child in v._children:
build_topo(child)
# TODO: v is finished only after all its children - add it to topo
build_topo(self)
# Step 2: the loss moves 1:1 with itself
self.grad = 1
# Step 3: walk the order BACKWARDS, so each node is complete before
# it hands its gradient on to its children.
for v in reversed(topo):
for child, local_grad in zip(v._children, v._local_grads):
pass # TODO: chain rule - child's grad grows by local_grad * v.grad
# Try it on the example from the video:
a, b, c = Value(2.0), Value(-3.0), Value(10.0)
e = a * b # -6
d = e + c # 4
L = d * e # -24 (e is used twice: by d and by L)
L.backward()
print('L =', L.data)
print('a.grad =', a.grad, ' b.grad =', b.grad, ' c.grad =', c.grad, ' e.grad =', e.grad)
4Check yourself
5Unlocked in microgpt
backward() (lines 59 to 72) is now complete: lines 60 to 68 are your topological sort from lesson 3.3, and lines 69 to 72 are the chain rule you just wrote. Line 172, loss.backward(), runs it once per training step, across about 60,000 nodes for a typical name. PyTorch and JAX (the tools real AI labs use to train big models) do exactly this, on big grids of numbers instead of single numbers.
Progress is saved in this browser. Sign in to keep it across devices.