JAX and Flax

Optax

Optax is JAX's optimizer library — each optimizer is a pair of pure functions plus a state you carry, and complex training recipes snap together from small pieces.

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.

Optax is the box of optimizers for JAX — the update rules that turn gradients into actual parameter improvements — built so that even the optimizer's memory is something you hold yourself.

Think of a fitness coach with a small notebook. You report today's effort (the gradient). The coach consults the notebook — recent momentum, how jumpy each exercise has been — writes today's entry, and prescribes the exact adjustment to make. The prescription depends on the history, not only on today.

In Optax, the coach is stateless between visits: the notebook is a separate thing called the optimizer state, and you carry it to every session. JAX's no-hidden-state rule, applied to the coach.

Why it exists

Plain gradient descent — step downhill, fixed stride — works but wastes effort. Modern optimizers like Adam adapt: they build momentum in directions that keep paying off, and tread carefully where gradients jump around. All of that adaptation is memory, and memory must live somewhere.

Frameworks with optimizer objects hide it inside (optimizer.step() mutates silently). JAX cannot — hidden mutation breaks the compiler's assumptions. So Optax makes the memory explicit. An optimizer is two pure functions: init (create the notebook) and update (read gradient + notebook, return prescription + updated notebook).

How it works

opt_state = init(params)                        blank notebook, one page per knob

each step:
  grads                    ──┐
  opt_state (notebook)     ──┼──▶  update  ──▶  updates + new notebook
                             │
  params + updates  ──▶  apply_updates  ──▶  new params

One more idea makes Optax special: chaining. Recipes snap together from small pieces. Take "clip extreme gradients, THEN apply Adam, THEN scale by a schedule". That is three links in a chain, each a tiny transformation of the gradient stream.

A real example you have seen

Cruise control in a car adapts like this. It does not react to this instant's speed alone. It tracks how the speed has been trending and eases the throttle accordingly, so the car neither jerks nor drifts. Adam is cruise control for learning: smooth where the road is smooth, cautious where it is bumpy.

Remember this

  • An Optax optimizer = init + update, both pure; the state is yours to carry.
  • update returns updates + new state; apply_updates applies them to params.
  • Recipes chain from small pieces: clip → adam → schedule.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install optax

This pulls in JAX. Outputs verified with jax 0.6.2 and optax 0.2.8, CPU.

Adam on the tiny regression, by hand

optax_loop.py
import jax
import jax.numpy as jnp
import optax

def loss_fn(params, x, y):
    pred = params["w"] * x + params["b"]
    return jnp.mean((pred - y) ** 2)


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

optimizer = optax.adam(learning_rate=0.1)
opt_state = optimizer.init(params)            # Adam's moving averages live here

for step in range(300):
    grads = jax.grad(loss_fn)(params, x, y)
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)

print(f"w={params['w']:.3f}, b={params['b']:.3f}")
print(f"loss={loss_fn(params, x, y):.6f}")
Output
w=1.989, b=1.030
loss=0.000156

The walkthrough

optimizer.init(params) builds state shaped like the parameters — for Adam, two moving-average trees (momentum and squared-gradient memory) mirroring your params leaf-for-leaf. Pytrees all the way down: print jax.tree.map(lambda a: a.shape, opt_state) and recognise every shape.

The three-line rhythm — grad, update, apply_updates — is the entire optimizer API. Notice what update returns: not new parameters but updates, the deltas to add. Keeping the delta separate is what lets further transformations process it before application.

Both new values are caught. updates, opt_state = ... — dropping the returned opt_state and reusing the old one silently downgrades Adam toward memoryless behaviour. Carrying state is your job now; this is the price of purity, and this line is where it is paid.

Chaining and schedules — the composability that makes Optax more than an Adam supplier:

python
schedule = optax.cosine_decay_schedule(init_value=0.1, decay_steps=300)
optimizer = optax.chain(
    optax.clip_by_global_norm(1.0),     # tame exploding gradients first
    optax.adam(learning_rate=schedule), # then Adam, with decaying rate
)

Same init/update interface; the chain runs each link over the gradient stream in order. Gradient clipping as a composable piece, not an optimizer feature.

Common mistakes

Discarding the new optimizer state. updates, _ = optimizer.update(...) compiles, runs, and trains worse than it should — the classic silent Optax bug. The state must flow forward every step.

Applying updates by hand. params = jax.tree.map(lambda p, u: p - u, params, updates) has a sign bug waiting: Optax updates are already negated — apply_updates adds them. Use apply_updates and never think about the sign again.

Re-initialising state every step. Calling optimizer.init(params) inside the loop wipes Adam's memory each iteration — training limps along like raw SGD with odd scaling. init runs once, before the loop.

Mismatched trees after model surgery. Add a layer to params but keep the old opt_state, and update fails with a structure mismatch. After any change to the parameter tree, re-init (or migrate) the state.

Try it yourself

Swap optax.adam(0.1) for optax.sgd(0.1) and rerun — compare the final loss at 300 steps. Then give SGD momentum=0.9 and watch it close most of the gap. You have reproduced, in one minute, the momentum story told in gradient descent.

What to learn next

Researcher — Mathematics and papers.

The GradientTransformation algebra

Optax's core abstraction is the pair $(\text{init}, \text{update})$ with signature $\text{update}: (g_t, s_t, \theta_t) \mapsto (u_t, s_{t+1})$ — symbols: $g_t$ the gradient tree at step $t$; $s_t$ the transformation state; $\theta_t$ the parameters (available for rules that need them, e.g. weight decay); $u_t$ the update tree, applied as $\theta_{t+1} = \theta_t + u_t$. Transformations compose by function composition of the $g$-stream with tupled state: chain(T1, T2) feeds T1's output as T2's input gradient. Adam itself decomposes as chain(scale_by_adam(), scale(-lr)) in Optax's own source — the named optimizers are prefab chains.

Adam's update (Kingma and Ba 2015): $m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t$, $v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2$, with bias corrections $\hat{m}_t = m_t/(1-\beta_1^t)$, $\hat{v}_t = v_t/(1-\beta_2^t)$, and $u_t = -\eta\, \hat{m}_t / (\sqrt{\hat{v}_t} + \epsilon)$. State cost: $2|\theta|$ extra floats — the tree-shaped notebooks seen in init. AdamW (Loshchilov and Hutter 2019, Decoupled weight decay regularization) appends decoupled decay — in Optax, one more link (add_decayed_weights) rather than a new optimizer, which is the algebra earning its keep.

Purity as an experimental instrument

Explicit state has research value beyond ideology. Optimizer state is checkpointable and inspectable as data (histogram $\hat{v}$ to find pathological curvature); state is swappable mid-run (optimizer schedules across phases); multi_transform routes different chains to different parameter-tree paths by predicate (frozen backbones, per-layer learning rates) — all without framework hooks. And because update is pure, the whole optimizer step fuses into the jitted train step, optimizer arithmetic included — there is no Python between loss and applied update at execution time.

Correctness subtleties preserved from the literature: bias correction's dependence on a step counter (in the state); clip_by_global_norm computing the norm across the entire tree, matching Pascanu et al. (2013) rather than per-leaf clipping; schedules as functions of the step count carried in state, not wall-clock side channels.

References

  • Kingma and Ba (2015), Adam: a method for stochastic optimization.
  • Loshchilov and Hutter (2019), Decoupled weight decay regularization.
  • DeepMind (2020–), Optax: composable gradient transformation and optimisation, in JAX — library and design documentation.

What to learn next