Automatic Differentiation
PyTorch Autograd Explained - In-depth Tutorial - YouTube 2.5. Automatic Differentiation — Dive into Deep Learning documentation
Automatic Differentiation in JAX
JAX is a library for high-performance numerical and scientific computing on modern hardware, including GPUs and TPUs. It extends NumPy with automatic differentiation and uses function transformations to compute gradients efficiently. These gradients are central to optimization problems and machine learning algorithms.
Automatic Differentiation (Autodiff)
Automatic differentiation evaluates derivatives by composing the derivative rules of elementary operations along a program's calculation, using the chain rule. Unlike symbolic differentiation, it need not build a simplified formula; unlike finite differences, it does not estimate a derivative from nearby function values. It differentiates the implemented computation, subject to floating-point error and the available rules. A discontinuity, integer decision, unsupported operation, or undefined derivative is not made differentiable by autodiff. Higher derivatives require sufficient differentiability and can be expensive. Review derivatives and the matrix shapes of linear maps before Jacobian products.
Forward and Reverse Mode Autodiff
For , the Jacobian has entries and shape . Forward mode propagates an input direction to . Reverse mode propagates output weights to . Basis-vector seeds recover individual columns or rows; neither mode is limited to a single coordinate seed.
The transpose comes from the multivariable chain rule, not a convention of JAX. Let and let be a scalar loss. Each input can affect the loss through every output, so
Thus reverse mode uses . The sum is also why contributions from separate paths must accumulate. This is the scalar chain rule applied to all intermediate coordinates.
A full Jacobian therefore takes roughly forward sweeps or reverse sweeps. Forward mode suits few inputs and many outputs; reverse mode suits many inputs and few outputs, especially a scalar loss. Reverse mode underlies backpropagation, but must retain or recompute intermediates for the backward sweep. These are work/memory trade-offs, not a universal speed ranking.
Differentiation Graph
A differentiation graph is a graphical representation of a function's computation, recording each operation from input to output, its intermediate variables, and their dependencies.
Automatic differentiation uses this graph to apply the chain rule efficiently and accurately:
- In forward mode, derivatives propagate in the direction of the graph.
- In reverse mode, gradients accumulate in the opposite direction.
Open full-size imageThis PyTorch example computes Q = 3a³ − b² elementwise. Read downward to follow the forward calculation and upward to trace gradient propagation. Labels such as PowBackward0 name the stored backward functions; “(2)” is the input tensor’s shape. The JAX code below uses a different function, but likewise combines elementary derivative rules to obtain the result.
How JAX Implements Autodiff
JAX supports both forward and reverse mode to compute gradients, Jacobians, and Hessians efficiently. Its main interfaces are:
-
jax.grad:- Computes a function's gradient.
- Uses reverse mode.
- Suits cases with far more input dimensions than output dimensions, such as a neural network's scalar loss.
-
jax.jvp(Jacobian-vector product):- Computes a Jacobian-vector product.
- Is the basic forward-mode operation.
- Computes , where is the Jacobian and is a vector.
-
jax.vjp(vector-Jacobian product):- Computes a vector-Jacobian product.
- Is the basic reverse-mode operation.
- Computes , where is a vector and is the Jacobian.
Example Usage
This example uses JAX to compute a function's gradient:
import jax
import jax.numpy as jnp
def f(x):
return jnp.sin(x) * jnp.cos(x)
# Obtain the gradient function of f
grad_f = jax.grad(f)
# Compute the gradient of f at x = 1.0
print(grad_f(1.0))
This code snippet computes the derivative of the function at the point .
Work through a shared input
Let , , and . At , a forward seed gives and , so . This is .
For reverse mode, seed . Addition sends . The multiplication contributes to and to ; sine contributes to . Contributions to the same input must be added, giving . The two routes through are why simply overwriting its gradient is wrong.
In the earlier JAX example, the expected derivative is . jax.grad expects a scalar output (shape (), not a length-one vector) and ordinarily differentiates floating-point inputs. For vector outputs, choose jax.jvp, jax.vjp, jax.jacfwd, or jax.jacrev according to the product or matrix needed. At a kink such as at zero, an implementation's chosen rule is not proof that a classical derivative exists.