# Optax: a gradient processing library for JAX, built from small composable pieces

> Optax is a JAX library of optimizers, losses and gradient transforms that you combine yourself. It suits researchers who want to build custom update rules, not teams looking for a batteries-included training framework.

**google-deepmind/optax** — Optax is a gradient processing and optimization library for JAX.

- Repository: https://github.com/google-deepmind/optax
- Website: https://optax.readthedocs.io
- Stars: 2,343 · Forks: 371
- Language: Python
- License: Apache-2.0
- Published: 2026-09-10 · Updated: 2026-09-10 · Language: en
- Canonical page: https://hysenlabs.com/projects/google-deepmind-optax

## The problem Optax solves: gradient processing as building blocks, not a trainer

Most deep learning libraries ship an optimizer as a fixed object attached to a model. Optax inverts that. The README states the goal plainly: provide "building blocks that can be easily recombined in custom ways", and favor "small composable building blocks" over larger abstractions. The library is a gradient processing toolkit, not a training framework.

That distinction decides who it is for. If you are a researcher who wants to test a new update rule, or who needs a schedule chained to a clipping transform chained to a masking transform, Optax gives you the primitives and expects you to assemble them. If you want a training loop, checkpointing and distributed execution handled for you, Optax is only one layer of that stack, and the README points to the broader DeepMind JAX Ecosystem rather than claiming to be the whole thing.

The project's own history supports the positioning. An initial prototype lived in JAX's experimental folder as jax.experimental.optix, was adopted widely inside DeepMind, and was then moved out of experimental, renamed optax and released as a standalone library. The pyproject.toml still classifies it as "Development Status :: 4 - Beta", which is worth noting for anyone treating API stability as a hard requirement.

## How optax.adam, optax.l2_loss and optax.apply_updates fit together

The mechanism is a small state machine around three functions. You construct an optimizer, for example optax.adam(learning_rate). You call optimizer.init(params) once to obtain an opt_state, a pytree holding whatever statistics that optimizer needs (for Adam, moment estimates). Then each step you compute gradients with jax.grad, pass them to optimizer.update(grads, opt_state), and receive two things: the updates to apply and the new opt_state. Finally optax.apply_updates(params, updates) adds the updates to the parameters.

What makes this composable is that optimizers are themselves transformations over gradients. Because init and update are the whole interface, you can chain transforms, and the README's description of the library as gradient processing rather than parameter updating is the reason apply_updates exists as a separate convenience utility instead of being folded into update. Loss functions live alongside the optimizers in the same package, so optax.l2_loss can be dropped directly into a differentiated loss without an extra dependency.

The README also notes an implementation preference: where reasonable, the code prioritizes readability and structuring that matches standard equations over code reuse. That is a deliberate trade. It makes individual optimizers easier to check against a paper, at the cost of some duplication across the codebase.

## Installing Optax with pip and running a first Adam update

The README gives two installation routes. The released version comes from PyPI:

```bash
pip install optax
```

If you need unreleased changes, the development version installs straight from GitHub:

```bash
pip install git+https://github.com/google-deepmind/optax.git
```

Either way, check the version floors in pyproject.toml before installing: Python 3.10 or newer, and jax>=0.5.3 plus jaxlib>=0.5.3. Optax declares absl-py and numpy as its other runtime dependencies, so a JAX installation is a prerequisite, not something Optax brings with it.

A first real use follows the README quickstart. You build the optimizer, initialize state from a parameter pytree, differentiate a loss, then update:

```python
optimizer = optax.adam(learning_rate)
params = {'w': jnp.ones((num_weights,))}
opt_state = optimizer.init(params)

compute_loss = lambda params, x, y: optax.l2_loss(params['w'].dot(x), y)
grads = jax.grad(compute_loss)(params, xs, ys)

updates, opt_state = optimizer.update(grads, opt_state)
params = optax.apply_updates(params, updates)
```

After the first update call, opt_state holds the optimizer's running statistics and params holds the adjusted weights. The README points to a getting-started notebook in docs/ for continuing past this point, and the examples/ directory contains notebooks for gradient accumulation, L-BFGS, meta-learning and freezing parameters, among others.

## Where Optax stops: no trainer, no data pipeline, no framework

The most common wrong turn is treating Optax as a training framework. It is not one. There is no data loader, no checkpointing, no mixed-precision policy, no distributed sharding strategy in the package as described. You supply the loop, the batching and the device placement; Optax handles the gradient-to-update step and nothing beyond it.

The dependency floor is a second practical constraint. Because Optax requires jax>=0.5.3 and jaxlib>=0.5.3, an environment pinned to an older JAX cannot simply add Optax on top. The pyproject.toml comment asks maintainers to keep jax and jaxlib versions in sync with the CI workflow and the conda-forge feedstock, which means the floor moves with JAX releases rather than staying fixed.

The Beta classifier is a third signal. It does not mean the code is unstable, and the repository is not archived, but it does mean the project does not promise API freeze. If you are writing code that must survive upgrades without edits, pin the version and read the release notes for v0.2.8, v0.2.7 and v0.2.6 before bumping.

Finally, Optax is JAX-only. If your model is written in another framework, the optimizers here are not directly usable; you would be reimplementing update rules rather than adopting this library.

## Optax compared with Flax and with writing your own update rule

Flax is the comparison that comes up most, and the difference is layer, not quality. Flax provides the model definition and training machinery around a JAX computation; Optax provides the optimizer and gradient transforms that such a loop calls. They are complementary, and the optax pyproject.toml lists flax>=0.5.3 in its test extras, which is a reasonable hint about how the two are exercised together. Choosing between them is a category error. Choosing whether you need both is the real question.

The sharper alternative is writing the update rule yourself. For a plain SGD step, a hand-written parameter update is a few lines and removes a dependency. What Optax adds is the state plumbing: the init/update pair handles per-parameter statistics, and transforms compose so that schedules, clipping and masking can be layered without you threading extra state through your loop by hand. That is the actual value, and it only materializes once your update rule is more complicated than gradient times learning rate.

For nonlinear solving rather than neural network training, the README points elsewhere, naming optimistix for root finding, minimisation, fixed points and least squares. That is a different problem class, and Optax is not the tool for it.

## Maintenance, versioning and the Apache-2.0 licence

The repository is not archived and the last push was on 2026-09-09, so the codebase is being touched. The release cadence visible in the release history is roughly two releases in early 2026 followed by a gap back to v0.2.6 in September 2025, which suggests releases cluster around batches of work rather than a fixed schedule. There is no published support window or long-term-support branch in the README.

The upgrade cost is mostly the JAX floor. Because Optax tracks jax and jaxlib versions, upgrading Optax can force a JAX upgrade, and a JAX upgrade can in turn affect the rest of your stack. The practical mitigation is to pin optax in your lockfile and read the release notes for the three most recent versions before moving.

On licensing, the project ships under Apache-2.0, and pyproject.toml points the license field at the LICENSE file rather than restating terms inline. Apache-2.0 is a permissive licence that includes an explicit patent grant, but the details matter and depend on how you distribute your own work. That is a question for your legal team, not for this article.

Contributions are open. The README asks that anyone adding a feature, such as a new optimizer, open an issue first, and the development section describes running sh test.sh for tests and make html -C docs for documentation after installing the docs extras.

## Conclusion

Adopt Optax if you already write JAX training loops and want optimizers, loss functions and gradient transforms you can recombine rather than a framework that owns the loop. Do not adopt it if you want a full training stack with checkpointing, data loading and distributed orchestration, or if your model is written in PyTorch, since Optax operates on JAX arrays. Before committing, verify that your installed JAX and jaxlib satisfy the jax>=0.5.3 and jaxlib>=0.5.3 floors declared in pyproject.toml, check that your Python is 3.10 or newer, and read the optimizers page of the documentation to confirm the specific update rule you need is implemented rather than something you would have to write with optax.GradientTransformation.

## FAQ

### how to install optax

Install the released version with pip install optax, or the development version with pip install git+https://github.com/google-deepmind/optax.git. Optax requires Python 3.10 or newer and jax>=0.5.3 with jaxlib>=0.5.3.

### what is optax

Optax is a gradient processing and optimization library for JAX, providing optimizers, loss functions and gradient transforms as small composable building blocks. It was originally prototyped in JAX's experimental folder as jax.experimental.optix.

### optax vs flax

They operate at different layers rather than competing. Flax provides model and training machinery for JAX, while Optax supplies the optimizer and gradient transforms that a training loop calls; flax appears in Optax's test extras in pyproject.toml.

### optax vs optimistic

Optax is a JAX gradient processing and optimization library. The README describes no component named optimistic and no comparison with anything by that name, so the two are not alternatives within this project's documentation.

## Sources

- [google-deepmind/optax on GitHub](https://github.com/google-deepmind/optax)
- [License: Apache-2.0](https://github.com/google-deepmind/optax/blob/main/LICENSE)
- [Project website](https://optax.readthedocs.io)
- [README](https://github.com/google-deepmind/optax/blob/main/README.md)
- [Releases](https://github.com/google-deepmind/optax/releases)

---

Hysen Labs editorial analysis, written from the project's own repository and release notes. Cite the canonical page: https://hysenlabs.com/projects/google-deepmind-optax
