JAX and Flax

grad and value_and_grad

jax.grad takes a function and returns a new function — its derivative — and value_and_grad returns both answer and slope in one pass.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

jax.grad is a machine that eats a function and spits out a new function — one that tells you the slope of the original at any point you ask.

Think of handing a recipe to a food scientist. She hands back a different document: for every ingredient, how strongly one extra pinch would change the taste. Salt: strong effect, add less. Water: barely any effect. That influence report — one number per ingredient — is the gradient.

The strange, wonderful part: you give grad a function, and it gives you back a function. Call the new one wherever you like and read off the slopes there.

Why it exists

All of deep learning runs on one question: which way should each knob turn to shrink the error? A model has millions of knobs. The gradient answers for all of them at once, and the training loop asks the question millions of times.

Other frameworks weave this ability into special objects — recording tapes, tracked variables. JAX keeps it as one honest function-to-function step. Your loss is a plain function; jax.grad(loss) is its slope; there is nothing else to learn. This mirrors how slopes are taught at school: work out the slope rule first, then read it off at a point. That is why many researchers find JAX the most natural of the frameworks.

How it works

   "square the number"   ──grad──▶   "the slope of the square"

   ask the first  about 3   ▶   9     the value there
   ask the second about 3   ▶   6     the slope there

   and it stacks:   grad of a grad   ▶   the slope of the slope

Because the output is an ordinary function, you can transform it again. A slope function of a slope function costs one more wrap.

A real example you have seen

Think of trekking in fog, or hill climbing in a mobile game map. You cannot see the peak, but you can feel the slope under your feet and step uphill. Every model that recommends, translates or completes text was trained by feeling slopes and stepping downhill on error, step after step. Those slopes come from exactly this kind of machinery. The walking itself is gradient descent.

Remember this

  • grad maps a function to its slope function — functions in, functions out.
  • The gradient = one influence number per input knob.
  • value_and_grad returns answer and slopes together — the training-loop workhorse.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install jax

Outputs verified with jax 0.6.2, CPU. Everything here is deterministic.

From one slope to a full loss gradient

grad_basics.py
import jax
import jax.numpy as jnp

# 1. grad turns a function into its derivative function
def f(x):
    return x ** 2


df = jax.grad(f)
print("f'(3) =", df(3.0))
print("f''(3) =", jax.grad(jax.grad(f))(3.0))     # derivatives compose

# 2. A real loss over parameters stored in a dict
def loss_fn(params, x, y):
    pred = params["w"] * x + params["b"]
    return jnp.mean((pred - y) ** 2)


params = {"w": 0.0, "b": 0.0}
x = jnp.array([1., 2., 3., 4.])
y = jnp.array([3., 5., 7., 9.])                    # truth: y = 2x + 1

loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
print("loss:", loss)
print("gradients:", grads)
Output
f'(3) = 6.0
f''(3) = 2.0
loss: 41.0
gradients: {'b': Array(-12., dtype=float32, weak_type=True), 'w': Array(-35., dtype=float32, weak_type=True)}

The walkthrough

Check the maths by hand first. Derivative of x² is 2x → 6 at x=3. Second derivative is the constant 2. When the paper answer and the machine agree, you can trust the machine on functions too big for paper.

grad differentiates with respect to argument 0 by default. loss_fn(params, x, y) therefore gets gradients for params only — x and y are treated as fixed data, which is what training wants. Different targets: jax.grad(loss_fn, argnums=1) or a tuple for several at once.

The gradient of a dict is a dict. Feed parameters as {"w": ..., "b": ...} and gradients come back in the same shape of container — matching leaf for leaf. This works for any nesting (dicts of lists of arrays), a structure called a pytree that gets its own lesson. No .grad attributes, no tape objects: containers in, matching containers out.

Sanity-check the signs. Loss 41 with all-zero parameters; gradient for w is −35, meaning increase w to reduce loss. True w is 2, current is 0 — the machine points the right way.

value_and_grad is one pass, not two. Computing the loss forward and the gradient backward shares the forward computation, so asking for both together is essentially free — and every training loop wants both (the loss for logging, the grads for stepping). Wrapping it in jit compiles the pair into one program.

Common mistakes

Non-scalar output. grad demands a function returning one number. A vector output raises: TypeError: Gradient only defined for scalar-output functions. Output had shape: (2,). Losses end in mean or sum for exactly this reason. Per-example derivative vectors are a different tool — Jacobians, via jax.jacobian or vmap of grad.

Differentiating through casts to Python. float(x), int(x), x.item() inside the function eject the value from JAX's world; grad errors or returns nonsense. Keep everything jnp until the end.

Integer inputs. Slopes need continuous inputs. jax.grad(f)(3) with a Python int raises a dtype error; write 3.0.

Expecting an accumulating .grad attribute. PyTorch habits die hard: JAX has no zero_grad, no accumulation, no attributes. Each call to the gradient function is a fresh, stateless evaluation. Coming from PyTorch autograd, this is the mental switch.

Try it yourself

Add a third parameter params["c"] with the model w*x + b + c*x**2, rerun, and read the gradient dict. Then take one manual step — params = {k: params[k] - 0.01 * grads[k] for k in params} — recompute the loss, and confirm it dropped below 41.

What to learn next

Researcher — Mathematics and papers.

The transformation view

grad is defined compositionally: trace the function to a jaxpr, then apply reverse-mode differentiation as a program transformation, producing a new jaxpr evaluated like any other. Formally, for $f: \mathbb{R}^n \to \mathbb{R}$, grad(f) computes $\nabla f(x) = \left(\partial f/\partial x_1, \dots, \partial f/\partial x_n\right)$ — symbols: $x_i$ the $i$-th input component; $\nabla f$ the vector of partial derivatives — in one reverse sweep costing a small constant times the forward evaluation, independent of $n$. The machinery underneath is the JVP/VJP pair: forward-mode jax.jvp computes Jacobian-vector products; reverse-mode is a transposition of the linearised program (jax.vjp), and grad(f)(x) is vjp applied to the unit cotangent. The design is spelled out in Frostig, Johnson, Leary (2018), Compiling machine learning programs via high-level tracing; the deeper theory of differentiation as program transformation goes back to Wengert (1964) and is surveyed in Baydin et al. (2018), Automatic differentiation in machine learning: a survey.

Why closure over composition matters

Because transformations return ordinary functions, they compose algebraically:

  • grad(grad(f)) — higher-order derivatives, arbitrarily deep (memory grows per level).
  • jit(grad(f)) — compiled gradients; the differentiated jaxpr fuses like any program.
  • vmap(grad(f)) — per-example gradients in one vectorised pass, the primitive behind differential privacy's per-sample clipping (Abadi et al. 2016, Deep learning with differential privacy) — awkward to express efficiently in tape-based frameworks.
  • jax.hessian(f) = jacfwd(jacrev(f)) — forward-over-reverse, the right mode ordering for dense Hessians.

This closure property — transformations of transformations — is JAX's core research affordance: meta-learning (differentiate through an optimisation), implicit differentiation, and physics-informed losses (derivatives inside the loss) all reduce to nesting.

Reverse-mode's memory bill remains: activations of the forward pass are stored for the backward sweep, $O(\text{depth})$; jax.checkpoint (rematerialisation) trades recompute for memory exactly as in the general theory.

weak_type in the output marks values born from Python scalars, allowing dtype promotion to stay NumPy-like under JAX's stricter float32 default — cosmetic here, occasionally load-bearing in dtype bugs.

References

  • Baydin et al. (2018), Automatic differentiation in machine learning: a survey.
  • Abadi et al. (2016), Deep learning with differential privacy — per-example gradients as a primitive.
  • Blondel et al. (2022), Efficient and modular implicit differentiation — differentiating through solvers in JAX.

What to learn next