Skip to main content

Automatic Differentiation in Deep Learning: How Frameworks Compute Gradients

Calculating read time…

Automatic differentiation (autodiff) is the technique deep learning frameworks use to compute exact gradients of a numeric function by mechanically applying the chain rule to a recorded sequence of elementary operations, rather than deriving a symbolic formula or approximating with finite differences. It's the machinery that turns "here's my loss function" into "here's exactly how to nudge every parameter to reduce it." 🔗

Every framework that trains neural networks — PyTorch, TensorFlow, JAX — is, underneath the layers and optimizers, an automatic differentiation engine wearing a deep learning costume. Misunderstanding how it works leads to some of the most confusing bugs in the field: gradients that are silently zero, memory that mysteriously grows across a training loop, or a `.backward()` call that crashes with an error about a graph having already been freed. Understanding the mechanism turns those from mysteries into two-minute fixes. ⚙️

Diagram showing a forward pass moving left to right through multiply, add, and loss nodes, dropping a local derivative tag at each stop, followed by a backward pass moving right to left multiplying each tag into a running gradient, ending at the final gradient stored for the input tensor

🔀 Quick Comparison: Forward-Mode vs. Reverse-Mode Autodiff

Property Forward-mode Reverse-mode
Pass direction Derivatives computed alongside the forward pass Values computed forward, derivatives computed in a separate backward pass
Efficient for Few inputs, many outputs Many inputs, few outputs (e.g., millions of parameters, one scalar loss)
Cost per pass One pass per input variable One pass computes gradients for all inputs at once
Memory pattern Low — no need to store the forward trace Higher — intermediate activations must be retained for the backward pass
Typical use in DL Jacobian-vector products (JAX's jvp) Standard neural network training (backpropagation)

1. Foundations: What Autodiff Actually Is

💭 Analogy first: imagine following a recipe and, at every single step (chop, sauté, simmer), writing down not just what you did but exactly how sensitive the final dish's taste is to a small change in that step. At the end, you can multiply all those little sensitivity notes together, in reverse order, to know exactly how much more salt at step 2 would change the final flavor — without ever re-cooking the whole dish. That's automatic differentiation.

There are three fundamentally different ways to get a derivative out of a computer, and it matters which one a framework uses. Numerical differentiation perturbs an input slightly and measures the change in output — simple, but slow and imprecise due to floating-point rounding error. Symbolic differentiation manipulates a mathematical expression algebraically to produce a new expression for the derivative — exact, but the resulting expression can explode in size for deeply nested functions. Automatic differentiation takes a third path: it decomposes a program into a sequence of elementary operations (add, multiply, exp, and so on) whose individual derivatives are known exactly, then mechanically composes them with the chain rule. The result is exact (to floating-point precision) and scales linearly with the number of operations in the original program, which is precisely why it is the only practical approach for functions with millions of parameters, like a neural network's loss.

🎯 Use this when you need to explain to a colleague why "autodiff" isn't the same thing as symbolic math or numerical estimation, and why that distinction is what makes training large models computationally feasible.

2. The Mechanism: Tracing and the Chain Rule

💭 Analogy first: a detective walking backward through a crime scene, retracing each footprint from the exit back to the entrance, learns the full path without needing a map drawn in advance — they just need to know, at each footprint, which direction the previous one came from. Reverse-mode autodiff retraces a computation the same way.

What it does: as your program runs, the framework records every operation performed on any tensor that needs a gradient, along with each operation's local derivative rule, into a directed acyclic graph (DAG). Why it's needed: a neural network's loss is a composition of potentially thousands of nested operations; the chain rule says the derivative of a composition is the product of the derivatives of its parts, but working that out by hand for a real model is not feasible, so the framework needs to do this bookkeeping automatically.

How it works, step by step:

  1. Forward pass: each operation computes its numeric output and stores a reference to its own local derivative function and its inputs (this is what PyTorch calls a grad_fn).
  2. A scalar output (typically the loss) is produced at the end of the forward pass.
  3. Backward pass begins with a seed gradient of 1.0 with respect to the loss itself.
  4. The graph is traversed in reverse topological order; at each node, the incoming gradient is multiplied by that node's local derivative and passed further back — exactly the chain rule, applied one recorded step at a time.
  5. When a leaf tensor (one with no further inputs, such as a model parameter) is reached, the accumulated gradient is stored — in PyTorch's case, added into that tensor's .grad attribute.

What fails without this record-and-replay approach: without a recorded graph, there is nothing to walk backward through, so the only remaining options are numerical differentiation (too slow and imprecise for millions of parameters) or hand-derived symbolic gradients (impractical to maintain as architectures change).

3. Forward-Mode vs. Reverse-Mode Differentiation

Both modes apply the same chain rule but propagate information in opposite directions, and the right choice depends entirely on the shape of your function — how many inputs versus how many outputs it has.

Reverse-mode (what "backpropagation" refers to in deep learning) computes the gradient of one scalar output with respect to every input in a single backward pass. This is exactly the shape of a training loss function — millions of parameters in, one scalar loss out — which is why it's the default in PyTorch's .backward(), TensorFlow's GradientTape.gradient(), and JAX's grad().

Forward-mode computes how one input's perturbation propagates forward to every output, in the same pass as the computation itself. It's efficient when there are few inputs and many outputs — the reverse of the typical training scenario — and it's documented as a first-class operation in JAX, which exposes it directly as jax.jvp (Jacobian-vector product) alongside reverse-mode's jax.vjp (vector-Jacobian product), and states that the two can be composed arbitrarily — including to compute full Hessian matrices by combining both modes.

What fails without the right mode: using reverse-mode when you have a huge number of outputs and few inputs (rare in typical DL, common in some scientific computing) wastes memory retaining a trace for outputs that didn't need it; using forward-mode for a standard training loss requires one pass per parameter, which is completely impractical at model scale.

💡 Trade-off: reverse-mode's efficiency for many-parameters/one-output problems comes at the cost of memory — every intermediate value needed to compute local derivatives during the backward pass must be kept around from the forward pass, which is exactly why very deep or very wide models can run out of memory even when the forward pass alone fits comfortably.

4. Real Example: PyTorch's Autograd

PyTorch's own tutorial documentation describes autograd's DAG directly: leaves are the input tensors, roots are the output tensors, and tracing the graph from roots to leaves lets the chain rule be applied automatically. During the forward pass, autograd simultaneously runs the requested operation and records that operation's gradient function into the DAG; the backward pass begins when .backward() is called on the root, at which point autograd computes gradients from each .grad_fn, accumulates them into each tensor's .grad attribute, and propagates all the way back to the leaves.

Two documented details matter enormously in practice. First, PyTorch's documentation explicitly notes the graph is dynamic — it is rebuilt from scratch after every forward pass, which is exactly what allows ordinary Python control flow (if-statements, loops with a variable number of iterations) inside a model's forward method; there is no separate "graph compilation" step to fight against. Second, calling .backward() a second time accumulates into the existing .grad rather than overwriting it — the documentation demonstrates this directly by showing the same call producing a larger cumulative gradient the second time it's run without clearing the previous value first.

✅ Worked example: for x = torch.tensor([2.0, 3.0], requires_grad=True) and z = (x**2 + 3*x).sum(), calling z.backward() populates x.grad with the analytically correct derivative of 2x + 3 evaluated at each element — no formula was ever written down by the programmer.

5. Real Example: TensorFlow's GradientTape

TensorFlow takes an explicit "recording" metaphor even further with tf.GradientTape. Its official documentation states that operations are recorded only if they execute inside the tape's context manager and at least one input is being "watched"; trainable tf.Variable objects are watched automatically, while plain tensors must be watched explicitly with tape.watch().

By default, the documentation states, the resources held by a tape are released as soon as .gradient() is called once — calling it a second time on the same tape raises an error. To compute multiple gradients from the same recorded computation (for example, gradients with respect to two different variables from the same forward pass), the documentation directs you to construct the tape with persistent=True. The same documentation also shows tapes can be nested to compute higher-order derivatives — an inner tape's gradient can itself be differentiated by an outer tape to obtain a second derivative — and that passing watch_accessed_variables=False gives fine-grained control over exactly which variables get tracked, rather than watching every trainable variable touched inside the block.

What fails without care here: the documentation itself warns that with watch_accessed_variables=False, a Keras layer's variables must already exist (the layer must already be built) before entering the tape, or that first training iteration will silently produce no gradients at all for that layer.

6. Real Example: JAX's Functional Transformations

JAX takes a distinctly different design stance: instead of an object that records a graph as a side effect of running your code (as in PyTorch and TensorFlow), JAX's own documentation frames grad as a pure function transformation — you pass it a Python function and it hands back a new function that computes the gradient, with no tape or context manager involved. This composes directly with JAX's other transformations: the documentation shows jit(grad(loss)) to get a compiled gradient function, and vmap layered on top of that to compute per-example gradients efficiently. The documentation also confirms you can differentiate to any order simply by nesting calls to grad, and that ordinary Python control flow (an if statement inside the differentiated function) works correctly because the function is simply re-traced for each branch actually taken.

What this buys you: because differentiation is just a function transformation rather than a stateful recording process, JAX's approach documents cleanly composable higher-order derivatives and vectorized per-example gradients without any special-casing — the same building blocks (grad, jit, vmap) combine to express what would otherwise be bespoke code.

7. Implementation Patterns

The following is an original, minimal illustration of the recording mechanism itself — not production code, and deliberately simplified to show the idea without any framework's actual internals.

class ExampleNode: """A minimal illustration of recording + chain rule replay.""" def __init__(self, value, parents=(), local_grads=()): self.value = value self.parents = parents # nodes this one was built from self.local_grads = local_grads # d(self)/d(each parent) self.grad = 0.0 def __mul__(self, other): out_value = self.value * other.value # local derivative of a product w.r.t. each factor return ExampleNode(out_value, parents=(self, other), local_grads=(other.value, self.value)) def backward(self, seed=1.0): self.grad += seed for parent, local_grad in zip(self.parents, self.local_grads): parent.backward(seed * local_grad) x = ExampleNode(2.0) w = ExampleNode(3.0) y = x * w # forward pass: records the multiply and its local grads y.backward() # backward pass: replays the chain rule print(x.grad) # 3.0 == dy/dx == w's value print(w.grad) # 2.0 == dy/dw == x's value

This is deliberately the same idea shown in the diagram above, written as runnable code: each operation records its parents and local derivatives at forward time, and backward() does nothing more than walk that record, multiplying as it goes. Real frameworks add device dispatch, in-place operation tracking, and heavily optimized kernels on top of exactly this idea — the core mechanism does not change.

8. Enterprise Rollout: Reliability and Governance for Gradients

Gradient correctness is invisible until it silently corrupts a training run, which makes this one of the areas where enterprise teams most need process, not just code review.

Gradient checking as a CI gate: whenever a team implements a custom operation with a hand-written backward function, that backward function should be checked against a numerical approximation before it ever reaches production training code. PyTorch documents a built-in utility, autograd.gradcheck, for exactly this purpose — comparing an implemented analytic gradient against a finite-difference approximation. Treat a failing gradcheck the same as a failing unit test: it blocks the merge.

Numerical stability monitoring: track gradient norms (not just loss) on a training dashboard. A gradient norm that explodes toward infinity or collapses toward zero over a run is often visible long before the loss curve shows anything unusual, and catching it early avoids burning GPU-hours on a training job that will need to be restarted from an earlier checkpoint anyway.

Mixed precision interactions: reduced-precision training changes the numeric range gradients can safely represent; a pipeline that switches precision modes without re-validating that gradients neither overflow nor underflow risks silent accuracy loss that won't show up as a crash.

Distributed gradient synchronization: in multi-device training, gradients computed locally on each device must be correctly aggregated (commonly averaged) across devices before the optimizer step. A CI gate that runs a small, deterministic multi-device job and checks that synchronized gradients match a known-good single-device computation catches synchronization bugs that are otherwise very hard to reproduce.

Ownership and versioning: when custom autograd functions or gradient-modifying utilities (gradient clipping thresholds, custom loss scaling) change, version that code alongside the model checkpoints trained with it, for the same reason data pipeline transforms need versioning — a "regression" that's actually a silent gradient-computation change is one of the hardest bugs to diagnose after the fact.

Access controls and incident response: production training infrastructure that exposes custom autograd extensions or modifies core training loop code should go through the same review gates as any other critical infrastructure change, with a documented rollback path if a change to gradient computation degrades a running job.

✅ Practical pattern: log gradient norm (and, ideally, the norm per layer) alongside loss on every training run's dashboard. When a run goes unstable, this is usually the first signal, well before the loss curve makes it obvious.

9. Common Mistakes

Forgetting to zero gradients between steps. Because gradients accumulate into .grad by design rather than being overwritten, skipping the zeroing step (optimizer.zero_grad() or equivalent) silently sums gradients across steps. Training still runs — the loss numbers even look plausible for a while — but the effective gradient magnitude grows every step, corrupting the optimization without ever raising an error.

Calling backward twice without retain_graph. By default, the graph built for a backward pass is freed once that pass completes, since it's normally rebuilt fresh on the next forward pass. Code that tries to call .backward() a second time on the same graph (for example, when computing two different losses from one shared forward pass) will hit a runtime error unless the first call passed retain_graph=True — and setting that everywhere out of habit quietly leaks memory across a long training run.

In-place operations breaking the graph. Modifying a tensor in place after it has been used in the forward pass (for example, an in-place ReLU applied to a tensor autograd still needs in its original form) can corrupt the values a backward function needs, sometimes producing incorrect gradients without any warning, and other times raising a version-counter error that can be confusing without understanding why the check exists.

Overusing detach() or forgetting no_grad(). Detaching a tensor that should still receive gradient flow silently zeroes out part of the intended training signal for that branch of the model — no error, just a part of the network that never learns. Conversely, forgetting to wrap inference-only code in a no-gradient context wastes memory building a graph that will never be used for a backward pass.

Watching the wrong tensor in TensorFlow. Since GradientTape only records operations on watched tensors, computing a loss from a plain (non-Variable) tensor that was never explicitly watched returns None for its gradient — as the framework's own documentation demonstrates directly with an unwatched variable — which is easy to misread as "the model isn't learning" rather than "the gradient was never computed at all."

❓ FAQ

Is automatic differentiation the same thing as backpropagation?

Not quite. Backpropagation is the specific application of reverse-mode automatic differentiation to neural networks. Automatic differentiation is the more general technique; reverse-mode is one of its two main strategies, and backpropagation is that strategy's name when used for training networks.

Why do gradients accumulate instead of being overwritten by default?

Accumulation supports use cases where a single parameter contributes to a loss through multiple paths in the same graph, or where you deliberately want to sum gradients from several backward calls (such as simulating a larger batch size across several smaller ones). The trade-off is that ordinary training loops must explicitly clear gradients before each new step.

Why does reverse-mode use more memory than forward-mode?

Reverse-mode needs the values computed during the forward pass to compute local derivatives during the backward pass, so those intermediate values must be kept in memory until the backward pass consumes them. Forward-mode computes derivatives alongside the forward pass itself and doesn't need to retain a separate trace afterward.

Can automatic differentiation handle if-statements and loops in my model code?

Yes, in frameworks that build the graph dynamically at run time, such as PyTorch's autograd or JAX's tracing-based transformations. Because the graph reflects the actual operations executed on a given call, whichever branch of an if-statement runs is exactly what gets differentiated — there's no separate static graph that would need to represent every possible branch in advance.

How do I verify a custom gradient implementation is correct?

Compare it against a numerical approximation obtained by slightly perturbing the input and measuring the change in output. This is exactly what dedicated gradient-checking utilities, such as PyTorch's documented gradcheck, automate for you.

🔗 References & Further Reading

PyTorch, TensorFlow, and JAX are trademarks of their respective owners; 

📝 Summary

  • Automatic differentiation computes exact gradients by decomposing a function into elementary operations and mechanically applying the chain rule.
  • It differs from numerical differentiation (approximate) and symbolic differentiation (exact but can explode in expression size).
  • Reverse-mode records a forward-pass trace, then walks it backward, multiplying local derivatives to build a full gradient in one pass — ideal for many-parameters, one-loss problems.
  • Forward-mode propagates derivatives alongside the forward computation and is efficient for few-inputs, many-outputs problems.
  • PyTorch's autograd builds a dynamic DAG per forward pass and accumulates gradients into leaf tensors' .grad attributes.
  • TensorFlow's GradientTape explicitly records watched operations and requires persistent=True for multiple gradient() calls on one trace.
  • JAX treats differentiation as a pure, composable function transformation rather than a stateful recording process.
  • Most real-world autodiff bugs are silent — unzeroed gradients, broken graphs from in-place ops, or unwatched tensors — rather than loud crashes.

Thanks for reading — may your gradients flow and your graphs stay intact. 🚀

Comments