AI tutorialsMighty Professional
Build a Language Model · Learning from data

Automatic Differentiation from Scratch

The neural networks page derived backpropagation by hand for one architecture. A framework derives it for any program, at a cost of a few times the forward pass, by recording what the program did and running the chain rule backwards over the record. This page builds that machinery, measures why reverse mode wins for training, shows what it costs in memory and how checkpointing buys that memory back, and covers the places where "automatic" still needs judgement.

Time~50 minLevelIntermediatePrereqsNeural Networks, especially backpropagation; Linear Models for gradient checks.StackC++20 & Rust · browser demos
◂ Build a Language ModelPhase 1 · Learning from dataNext · Training deep networks ▸

01A program is a graph of local derivatives

Every numerical program decomposes into primitive operations, each with a derivative that is easy to write down: the derivative of a sum is 1 toward each input, of a product is the other input, of sin is cos. Recording which operation consumed which value gives a , and the derivative of the output with respect to any input is a sum over every path from that input to the output of the product of the local derivatives along the path.[1]

∂f/∂x = Σpaths Πedges ∂child/∂parent

Enumerating paths is exponential in the depth of the graph. Forward mode and reverse mode compute the same sum by merging paths at every node, and each touches every edge exactly once.[1]

The widget below records f(a, b) = (a·b + sin a)·b as six nodes. The forward pass fills in values. The backward pass starts at the output with adjoint 1 and walks the nodes in reverse order of creation; each node multiplies its adjoint by its local derivatives and adds the products into its parents' adjoints. Reverse creation order is a valid topological order because a node is always created after the values it consumes.

Step a backward pass through a six-node graph

Each node shows its operation, its forward value, and ȳ, its adjoint so far. Both b and a feed two nodes each, so their adjoints receive two contributions and are complete only after both children have been processed. When all six steps are done the readout compares ∂f/∂a and ∂f/∂b with centered finite differences; they agree to the printed four decimals at the default inputs.

Three local-derivative patterns cover most of a network. An add node distributes its adjoint unchanged to every input. A max node routes it to the input that won the forward pass and sends zero to the others. A multiply node hands each input the adjoint scaled by the other input's value, which is why a weight multiplied by a large input receives a large gradient and why input scaling changes the effective learning rate.[7]

02Forward mode, reverse mode, and what each costs

carries a tangent through the graph in the same order as the values, so one pass yields the derivative of every node with respect to one input. carries adjoints backwards, so one pass yields the derivative of one output with respect to every node.[1] For a function with n inputs and m outputs, the full Jacobian costs n forward passes or m reverse passes.[2]

ċ = Σp (∂c/∂p) ṗ    p̄ = Σc c̄ (∂c/∂p)

The two rules are the same multiplication read from opposite ends. A forward pass computes a Jacobian-vector product; a reverse pass computes a vector-Jacobian product. Each costs a small constant times one evaluation of the function, about three times in JAX's accounting.[2]

A training loss maps millions of parameters to one number: n is enormous and m is 1. Reverse mode gets the whole gradient in a single pass; forward mode would need one pass per parameter.[2] That asymmetry is why reverse mode, under the name backpropagation, is the default for training; for a function with few inputs and many outputs, forward mode wins by the same margin.[1]

Count the work for a full Jacobian both ways

The graph is a seeded two-layer scalar network with 16 hidden units, built on the same tape as the other widgets. The widget runs n forward passes and m reverse passes and counts the multiply-adds each performs; nothing is estimated from a formula. At 32 inputs and 1 output, the loss-function shape, forward mode does 32 times the work of reverse mode. Slide the outputs up and the inputs down and the ratio inverts. One Jacobian entry is computed both ways and matches.

03Reverse mode in detail: adjoints accumulate

A value used in two places has two paths to the output, and its adjoint is the sum of both contributions. A reverse-mode engine therefore adds into a node's adjoint rather than assigning it.[7] PyTorch carries the same accumulation across calls: each backward pass adds into the .grad of leaf tensors instead of replacing it, which is why training loops zero the gradients before each backward pass.[3] Replacing the += with = is a silent bug: every number still comes out finite, and the gradient is simply wrong for any reused value.

Reuse a value, then break the accumulation

y = relu(x·w)·w + x·w uses x twice and w three times. With accumulation the tape agrees with finite differences to about 10⁻⁹. With overwriting, the last child to be processed wins and the earlier contributions are lost; at the defaults ∂y/∂w comes out as 0.96 instead of 3.12 and ∂y/∂x as 0.64 instead of 1.44, and the error changes with the inputs because it depends on which contribution survives.

Frameworks build the graph as the program runs, so each training iteration records a fresh graph; a Python if statement or loop that changes shape from one batch to the next is differentiated as executed.[3] A scalar engine of this kind, with a dynamically built graph and reverse-order backward, fits in about a hundred lines, which is the size of micrograd.[6] Production engines differ by operating on tensors, so that one node is a whole matrix multiply with a hand-written derivative rule, not a million scalar nodes.

04Memory, and buying it back with recomputation

Forward mode stores nothing: tangents are consumed as they are produced, so its memory does not depend on the depth of the computation. Reverse mode must keep the forward values that the local derivatives need (the input values of every multiply, the cosines of every sin) until the backward pass reaches them, so its memory grows with depth.[2] For a deep network this is the activation memory, and at large batch sizes or long sequences it can exceed the memory of the parameters themselves.

trades compute for that memory: the forward pass keeps only the inputs of selected segments and discards everything inside them, and the backward pass recomputes a segment's interior from its saved input when it gets there.[4][5] For a chain of L equal layers cut into segments of s, the forward pass stores about L/s boundary activations, and during the backward pass one segment's s interior activations are live at a time; the peak is near L/s + s, which is smallest at s = √L, where the peak is about 2√L in place of L. The price is one extra forward evaluation of every recomputed layer.

Choose a segment length for a chain of layers

The count treats every layer's activation as one unit. Without checkpointing all 64 activations are live at the start of the backward pass. With segments of 8, the forward pass stores 8 boundaries and the backward pass recomputes 56 layers, 88% extra forward work, for a peak of 16. The sweep bottoms out at s = 8 = √64. The count assumes each layer is recomputed once, from the nearest stored boundary; PyTorch's checkpoint_sequential instead reruns every segment except the last in full, which costs nearly one extra forward pass. Real policies are less uniform: JAX lets a checkpoint policy save compute-heavy results such as matrix products and recompute cheap ones.[5]

05Kinks, conventions and gradient checks

Autodiff is exact only where every primitive is differentiable. ReLU at 0 and sqrt at 0 are not, and a framework must still return something. PyTorch's documented rules, applied in order, are to use the derivative where it exists, otherwise a minimum-norm subgradient for locally convex functions (a minimum-norm supergradient for locally concave ones), otherwise a value defined by continuity; for ReLU at exactly 0 that gives 0.[3] The choice almost never matters for training, because landing exactly on a kink in floating point is rare, but it does mean the "gradient" of a piecewise-linear network is a convention at its corners.

A finite-difference check sees kinks differently: the centered difference averages the slopes on both sides of any kink within h of the test point, so it disagrees with autodiff there even when both are correct. Test gradients at generic points, and when a check fails, move the point before concluding the derivative rule is wrong.

Compare autodiff and finite differences across two kinks

g has kinks at w = 0 and w = 1. The green autodiff derivative comes from a reverse pass on the tape at each plotted point. Away from the kinks the yellow finite-difference curve sits on it. Within h of a kink the finite difference reports an average slope, and exactly at w = 1 autodiff reports 0 by the convention above while the finite difference reports 0.5. Shrinking h narrows the disagreement to a thinner band; it never removes the point itself.

06The engine trains a network

The scalar tape is enough to train a network. The widget builds a 2-input, 2-hidden-unit, 1-output ReLU network on the tape, one node per scalar operation, evaluates the squared-error loss over the four XOR cases, runs the backward pass, and steps the nine parameters. Every gradient the optimizer uses is produced by the same reverse pass as in §1, and the readout verifies it against finite differences at the current weights on every step.

Train XOR with the tape and check every gradient

Each Step applies ten gradient updates from a fixed starting point. The loss falls from 0.52 to below 10⁻³ within the first hundred updates at the default rate, and the output map develops the diagonal band that XOR needs: high at (0, 1) and (1, 0), low at the other two corners. Each loss evaluation builds a tape of 89 nodes. The gradient check reads about 10⁻¹² for the first ten Steps but 2.5·10⁻² at Step 0, where the (1, 1) case puts the first hidden unit's input at exactly 0.6 − 0.7 + 0.1 = 0, a kink. The trained network parks the (0, 0) and (1, 1) pre-activations within h of zero, so from Step 11 on the check settles near 3·10⁻⁶: the kink effect from §5, with both derivatives correct.

07A tape-based autograd, checked

These complete programs implement reverse-mode autodiff on a tape: nodes are appended in creation order with their value, their parents and the local derivatives toward each parent, and the backward pass walks the tape in reverse accumulating adjoints. The checks compare the gradient of the §1 expression and the §3 fan-out expression with centered finite differences, then run a small gradient-descent loop on the tape and assert it converges.

Scalar reverse-mode tape, complete programs
#include <cassert>
#include <cmath>
#include <vector>
struct Node {
    double value;
    double adjoint;        // ∂output/∂this, filled in by backward().
    int parents[2];        // Tape indices of the inputs; always smaller than this node's index.
    double locals[2];      // ∂this/∂parent for each parent, computed at forward time.
    int parentCount;
};
struct Tape {
    std::vector<Node> nodes;
    int record(double value, int parentA, int parentB, double localA, double localB, int parentCount) {
        nodes.push_back({value, 0.0, {parentA, parentB}, {localA, localB}, parentCount});
        return static_cast<int>(nodes.size()) - 1;
    }
    int leaf(double value)      { return record(value, -1, -1, 0, 0, 0); }
    int add(int a, int b)        { return record(nodes[a].value + nodes[b].value, a, b, 1, 1, 2); }
    int mul(int a, int b)        { return record(nodes[a].value * nodes[b].value, a, b, nodes[b].value, nodes[a].value, 2); }
    int sine(int a)              { return record(std::sin(nodes[a].value), a, -1, std::cos(nodes[a].value), 0, 1); }
    int relu(int a)              { const bool active = nodes[a].value > 0; return record(active ? nodes[a].value : 0, a, -1, active ? 1 : 0, 0, 1); }
    int square(int a)            { return record(nodes[a].value * nodes[a].value, a, -1, 2 * nodes[a].value, 0, 1); }
    // Reverse pass: seed the output with 1, then visit nodes in reverse creation order, which is a
    // topological order, and accumulate (never assign) each node's contribution into its parents.
    void backward(int output) {
        for (Node& node : nodes) node.adjoint = 0;
        nodes[output].adjoint = 1;
        for (int index = output; index >= 0; --index)
            for (int slot = 0; slot < nodes[index].parentCount; ++slot)
                nodes[nodes[index].parents[slot]].adjoint += nodes[index].locals[slot] * nodes[index].adjoint;
    }
};
// f(a, b) = (a·b + sin a)·b, built fresh on a tape each call; returns the output index.
int buildExpression(Tape& tape, int& a, int& b, double valueA, double valueB) {
    a = tape.leaf(valueA); b = tape.leaf(valueB);
    return tape.mul(tape.add(tape.mul(a, b), tape.sine(a)), b);
}
double expression(double a, double b) { return (a * b + std::sin(a)) * b; }
double fanOut(double x, double w) { return std::fmax(0.0, x * w) * w + x * w; }
int main() {
    const double step = 1e-6;
    Tape tape; int a, b;
    const int output = buildExpression(tape, a, b, 1.5, -2.0);
    tape.backward(output);
    const double numericalA = (expression(1.5 + step, -2.0) - expression(1.5 - step, -2.0)) / (2 * step);
    const double numericalB = (expression(1.5, -2.0 + step) - expression(1.5, -2.0 - step)) / (2 * step);
    assert(std::abs(tape.nodes[a].adjoint - numericalA) < 1e-6 && std::abs(tape.nodes[b].adjoint - numericalB) < 1e-6);
    Tape reuse;
    const int x = reuse.leaf(1.2), w = reuse.leaf(0.8);
    const int y = reuse.add(reuse.mul(reuse.relu(reuse.mul(x, w)), w), reuse.mul(x, w)); // w is used three times.
    reuse.backward(y);
    const double numericalW = (fanOut(1.2, 0.8 + step) - fanOut(1.2, 0.8 - step)) / (2 * step);
    assert(std::abs(reuse.nodes[w].adjoint - numericalW) < 1e-6); // Accumulation handles the fan-out.
    double weight = 0.0; // Minimize (3·weight − 6)² by gradient descent on the tape; the optimum is weight = 2.
    for (int iteration = 0; iteration < 200; ++iteration) {
        Tape loop;
        const int parameter = loop.leaf(weight);
        const int loss = loop.square(loop.add(loop.mul(loop.leaf(3.0), parameter), loop.leaf(-6.0)));
        loop.backward(loss);
        weight -= 0.05 * loop.nodes[parameter].adjoint; // A fresh tape per iteration: the graph is rebuilt as the program runs.
    }
    assert(std::abs(weight - 2.0) < 1e-9);
}
struct Node {
    value: f64,
    adjoint: f64,        // ∂output/∂this, filled in by backward().
    parents: [usize; 2],  // Tape indices of the inputs; always smaller than this node's index.
    locals: [f64; 2],     // ∂this/∂parent for each parent, computed at forward time.
    parent_count: usize,
}
struct Tape { nodes: Vec<Node> }
impl Tape {
    fn new() -> Tape { Tape { nodes: Vec::new() } }
    fn record(&mut self, value: f64, parents: [usize; 2], locals: [f64; 2], parent_count: usize) -> usize {
        self.nodes.push(Node { value, adjoint: 0.0, parents, locals, parent_count });
        self.nodes.len() - 1
    }
    fn leaf(&mut self, value: f64) -> usize { self.record(value, [0, 0], [0.0, 0.0], 0) }
    fn add(&mut self, a: usize, b: usize) -> usize { let value = self.nodes[a].value + self.nodes[b].value; self.record(value, [a, b], [1.0, 1.0], 2) }
    fn mul(&mut self, a: usize, b: usize) -> usize {
        let (value_a, value_b) = (self.nodes[a].value, self.nodes[b].value);
        self.record(value_a * value_b, [a, b], [value_b, value_a], 2)
    }
    fn sine(&mut self, a: usize) -> usize { let value = self.nodes[a].value; self.record(value.sin(), [a, 0], [value.cos(), 0.0], 1) }
    fn relu(&mut self, a: usize) -> usize {
        let value = self.nodes[a].value;
        let active = value > 0.0;
        self.record(if active { value } else { 0.0 }, [a, 0], [if active { 1.0 } else { 0.0 }, 0.0], 1)
    }
    fn square(&mut self, a: usize) -> usize { let value = self.nodes[a].value; self.record(value * value, [a, 0], [2.0 * value, 0.0], 1) }
    // Reverse pass: seed the output with 1, then visit nodes in reverse creation order, which is a
    // topological order, and accumulate (never assign) each node's contribution into its parents.
    fn backward(&mut self, output: usize) {
        for node in &mut self.nodes { node.adjoint = 0.0; }
        self.nodes[output].adjoint = 1.0;
        for index in (0..=output).rev() {
            for slot in 0..self.nodes[index].parent_count {
                let contribution = self.nodes[index].locals[slot] * self.nodes[index].adjoint;
                let parent = self.nodes[index].parents[slot];
                self.nodes[parent].adjoint += contribution;
            }
        }
    }
}
fn expression(a: f64, b: f64) -> f64 { (a * b + a.sin()) * b }
fn fan_out(x: f64, w: f64) -> f64 { (x * w).max(0.0) * w + x * w }
fn main() {
    let step = 1e-6;
    let mut tape = Tape::new(); // f(a, b) = (a·b + sin a)·b, built as the program runs.
    let (a, b) = (tape.leaf(1.5), tape.leaf(-2.0));
    let product = tape.mul(a, b);
    let sine = tape.sine(a);
    let sum = tape.add(product, sine);
    let output = tape.mul(sum, b);
    tape.backward(output);
    let numerical_a = (expression(1.5 + step, -2.0) - expression(1.5 - step, -2.0)) / (2.0 * step);
    let numerical_b = (expression(1.5, -2.0 + step) - expression(1.5, -2.0 - step)) / (2.0 * step);
    assert!((tape.nodes[a].adjoint - numerical_a).abs() < 1e-6 && (tape.nodes[b].adjoint - numerical_b).abs() < 1e-6);
    let mut reuse = Tape::new();
    let (x, w) = (reuse.leaf(1.2), reuse.leaf(0.8));
    let inner = reuse.mul(x, w);
    let gated = reuse.relu(inner);
    let left = reuse.mul(gated, w);
    let right = reuse.mul(x, w);
    let y = reuse.add(left, right); // w is used three times.
    reuse.backward(y);
    let numerical_w = (fan_out(1.2, 0.8 + step) - fan_out(1.2, 0.8 - step)) / (2.0 * step);
    assert!((reuse.nodes[w].adjoint - numerical_w).abs() < 1e-6); // Accumulation handles the fan-out.
    let mut weight = 0.0; // Minimize (3·weight − 6)² by gradient descent on the tape; the optimum is weight = 2.
    for _ in 0..200 {
        let mut tape = Tape::new(); // A fresh tape per iteration: the graph is rebuilt as the program runs.
        let parameter = tape.leaf(weight);
        let three = tape.leaf(3.0);
        let scaled = tape.mul(three, parameter);
        let minus_six = tape.leaf(-6.0);
        let residual = tape.add(scaled, minus_six);
        let loss = tape.square(residual);
        tape.backward(loss);
        weight -= 0.05 * tape.nodes[parameter].adjoint;
    }
    assert!((weight - 2.0).abs() < 1e-9);
}
What's intentionally missing

Tensor-valued nodes with matrix derivative rules, broadcasting, gradient-of-gradient (the tape would need to record the backward pass itself), saved-tensor release after use, checkpointing, in-place operation tracking, and any parallelism. A production engine spends most of its code on those; the reverse walk at its center is the loop above.

08Where autograd bites

Zero the gradients, every step

Because PyTorch-style engines accumulate into a parameter's gradient buffer across backward calls, the buffer still holds the previous step's value unless it is cleared. Forgetting to clear it adds the gradients of every past batch into the current update; the loss usually still falls for a while, which hides the bug.

In-place modification of a value the backward pass needs produces a wrong gradient or an error. PyTorch tracks a version counter on saved tensors and raises if one was overwritten; it discourages in-place operations in most code for exactly this reason.[3]

Checkpointed code must be deterministic between the forward and the recomputation. If the recomputed forward differs, through a global flag that changed or a dropout mask drawn from a different random state, the result is an error or a silently wrong gradient. PyTorch's checkpoint stashes and restores the random-number state by default for this reason.[4]

Numerical gradient checks cost one function evaluation per parameter and should test derivative rules, not replace them. Run them on a small model at generic points with h near 10⁻⁵ in double precision, as on the linear models page; in single precision the rounding floor is much higher and the check is correspondingly looser.

09What's next

With gradients available for any architecture, Training Deep Networks asks why deep stacks of layers are still hard to train: how initialization and normalization keep signals from vanishing or exploding across depth, what Adam and learning-rate schedules change about the descent, how mixed precision keeps the arithmetic fast without losing gradients, and what the scaling laws say about spending a compute budget.

10Sources

The cited documents describe the algorithms and the behaviour of two production engines. Widget readouts report counts and derivatives computed in the page by the tape described above.

  1. Christopher Olah, 2015. Calculus on Computational Graphs: Backpropagation. Derivatives as sums over paths, forward and reverse mode as factorizations that touch each edge once, which mode suits many inputs versus many outputs, and the repeated reinvention of reverse mode across fields (citing Griewank, 2010).
  2. JAX authors. The Autodiff Cookbook. Jacobian-vector and vector-Jacobian products, the roughly 3× per-pass cost, forward mode's depth-independent memory, reverse mode's memory growing with depth, and forward mode for tall versus reverse mode for wide Jacobians.
  3. PyTorch contributors. Autograd mechanics. The DAG recorded during the forward pass and rebuilt every iteration, gradient accumulation into leaf tensors, the rules for gradients at non-differentiable points, and the hazards of in-place operations.
  4. PyTorch contributors. torch.utils.checkpoint. Activation checkpointing as a trade of compute for memory, recomputation by re-invoking the function during backward, and the requirement that the recomputation match the forward pass.
  5. JAX authors. Gradient checkpointing with jax.checkpoint (jax.remat). Saved residuals, the memory-versus-FLOPs trade-off, and rematerialization policies that save selected results such as matrix products.
  6. Andrej Karpathy, 2020. micrograd. A scalar reverse-mode engine over a dynamically built DAG in about a hundred lines, with the backward pass in reverse topological order.
  7. Stanford CS231n course staff. Backpropagation, Intuitions. The add, max and multiply patterns in backward flow, the effect of input scale on weight gradients, and staged computation.