C.W.K.
Stream
Lesson 02 of 05 · published

Optax: Composable Gradient Transformations

~8 min · training, jax, tutorial

Level 0Curious
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

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)

Code

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
)
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)
# 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)

External links

Exercise

Compose AdamW + clip_by_global_norm + ema using optax.chain. Run one optimizer step. Confirm the state structure. The 'gradient transformations as Lego' philosophy is in this 4-line composition.

Progress

Progress is local-only — sign in to sync across devices.
Spotted a bug or have feedback on this page?Report an Issue

Comments 0

🔔 Reply notifications (sign in)
Sign inPlease sign in to comment.

No comments yet — be the first.