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.
- 7 min read
- 3 reading levels
- Published
Read these first
On this page 5
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 paramsOne 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. updatereturns updates + new state;apply_updatesapplies them to params.- Recipes chain from small pieces: clip → adam → schedule.
What to learn next
- A full training loop in JAX — the assembly: Flax + Optax + jit, end to end.
- Gradient descent — the base algorithm every chain elaborates.
- Gradient clipping — the most-used chain link, in depth.
Developer — Code and libraries.
Setup
pip install optaxThis pulls in JAX. Outputs verified with jax 0.6.2 and optax 0.2.8, CPU.
Adam on the tiny regression, by hand
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}")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:
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
- A full training loop in JAX — the assembly: Flax + Optax + jit, end to end.
- Gradient descent — the base algorithm every chain elaborates.
- Gradient clipping — the most-used chain link, in depth.
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
- A full training loop in JAX — the assembly: Flax + Optax + jit, end to end.
- Gradient descent — the base algorithm every chain elaborates.
- Gradient clipping — the most-used chain link, in depth.