Optax is JAX's gradient processing library. It doesn't just provide optimizers — it provides composable gradient transformations that you can chain together like LEGO bricks. An "optimizer" in Optax is really just a chain of transformations applied to gradients.
import optax
# Common optimizers — each is a gradient transformation
optimizer = optax.adam(learning_rate=1e-3)
optimizer = optax.adamw(learning_rate=1e-3, weight_decay=0.01)
optimizer = optax.sgd(learning_rate=0.1, momentum=0.9)
optimizer = optax.lion(learning_rate=1e-4) # newer optimizer
# But the real power is composition
optimizer = optax.chain(
optax.clip_by_global_norm(1.0), # gradient clipping
optax.adam(learning_rate=1e-3), # Adam optimizer
)
# Or even more custom:
optimizer = optax.chain(
optax.clip_by_global_norm(1.0), # clip gradients
optax.scale_by_adam(), # Adam scaling (no LR)
optax.add_decayed_weights(0.01), # L2 regularization
optax.scale(-1e-3), # apply learning rate
)
The Update Cycle
Optax separates initialization from updates, keeping everything functional:
import jax
import jax.numpy as jnp
import optax
# 1. Create optimizer
optimizer = optax.adamw(learning_rate=1e-3)
# 2. Initialize optimizer state from params
params = {'w': jnp.ones((3, 4)), 'b': jnp.zeros(4)}
opt_state = optimizer.init(params)
# 3. Compute gradients (however you like)
grads = jax.grad(loss_fn)(params, x, y)
# 4. Get updates from optimizer
updates, new_opt_state = optimizer.update(grads, opt_state, params)
# 5. Apply updates to params
new_params = optax.apply_updates(params, updates)
# The full cycle:
# grads → optimizer.update(grads, opt_state, params) → updates, new_opt_state
# ↓
# new_params = optax.apply_updates(params, updates)
⚠️ Pure Function Check
Notice that optimizer.update returns a new opt_state — it doesn't mutate the old one. And apply_updates returns new params — it doesn't mutate in place. Everything is functional. If you forget to use the returned values, your model won't learn (a common bug!).
Compare with PyTorch:
# PyTorch # JAX + Optax
# optimizer = Adam(model.params()) # optimizer = optax.adam(1e-3)
# # opt_state = optimizer.init(params)
# optimizer.zero_grad() # (not needed — grads are values)
# loss.backward() # grads = jax.grad(loss_fn)(params)
# optimizer.step() # updates, opt_state = optimizer.update(...)
# # params = optax.apply_updates(params, updates)