Every neural network learns the same way. It makes a guess, measures how wrong it was, works out how much each weight was to blame, and nudges each one a little. The blame step is backpropagation, and in PyTorch it's a single line: loss.backward().
I wanted to see what that line actually does, so I wrote it from scratch in TypeScript, following the idea behind Andrej Karpathy's micrograd. The engine is about 120 lines with no libraries, and it's what is running in your browser below.
Watch it learn
Spirals are a hard case on purpose. No straight line splits them, so the network has to bend its boundary around each arm. Every step is the same loop: guess every point, measure the loss, run backward(), nudge every weight.
- Step
- 0 / -
- Loss
- -
- Accuracy
- -
Where the gradients come from
Every number in the network is a value that remembers how it was made. Here's one neuron. Move the sliders and every gradient updates. Select any box to see the sum that produced its gradient, or step through backward() to watch it work from the output back to the inputs.
w₁ = 0.80. w₁ was multiplied by x₁, so its gradient is x₁ × (gradient of x₁w₁) = 1.50 × 0.92 = 1.37.
How backward() works
It's the chain rule, applied in reverse order. Each operation only knows its own local rule. For multiplication, each side's gradient is the other side's value times the output's gradient:
mul(other: Value | number): Value {
const o = Value.from(other);
const out = new Value(this.data * o.data, [this, o], '*');
out.backwardFn = () => {
this.grad += o.data * out.grad;
o.grad += this.data * out.grad;
};
return out;
}backward() sorts the graph so every value comes after the values it was made from, sets the output's gradient to 1, then walks the list backwards calling each of these. Gradients add up with += because a value used in two places takes blame from both. Forgetting that is the classic bug.
How I checked it
Every operation is tested against a finite difference: nudge the input by a millionth, see how far the output moves, and compare that with what backward() says. The same check runs on the full network's loss. All 12 tests pass, including one that trains a network on the moons data and checks it gets above 95% accuracy.