JAXopt: a discontinued JAX optimizer library, and what to do about it
Hardware accelerated, batchable and differentiable optimizers in JAX.
At a glance
- What is it?
- JAXopt packed differentiable, batchable solvers for constrained, root-finding and bi-level problems into JAX. The README now says it is no longer maintained, so the practical question is which parts moved to optax and which you still need.
- Who is it for?
- Adopt JAXopt only if you need implicit differentiation or a solver wrapper that optax does not provide, and only with the expectation that the repository will not receive fixes; the README states it is no longer maintained nor developed. Do not adopt it for plain neural-network training, where optax is the intended replacement.
- Can I use it commercially?
- Yes. Apache-2.0 is a permissive licence: you can use, modify and sell software built on it, as long as you keep its copyright and licence notices.
- Is it still maintained?
- Yes. The repository last received commits 23 days ago.
- What is it written in?
- Mainly Python, according to GitHub's language statistics.
Answers come from the project's GitHub data, last synced on October 1, 2026, and from our analysis. They are not legal advice.
Editorial analysis
What JAXopt was built for, and who still needs it
JAXopt targets a class of problems that ordinary gradient-descent libraries do not touch: constrained optimization, root finding, fixed points, and bi-level problems where one optimization sits inside another. The README frames the library around three properties. Implementations run on GPU and TPU in addition to CPU. Multiple instances of the same problem can be vectorized with JAX's vmap. And solutions can be differentiated with respect to their inputs, either implicitly or by autodiff of unrolled iterations.
That third property is the reason the library existed. If you are training a network, you differentiate a loss and hand the gradient to an update rule. If you are solving a linear program, a projection, or an equilibrium, the thing you want gradients through is the solver itself. JAXopt's implicit differentiation framework is described in a paper cited in the README, "Efficient and Modular Implicit Differentiation," and the repository carries an examples/implicit_diff directory alongside examples/constrained, examples/fixed_point and examples/deep_learning. The audience is researchers and engineers doing differentiable programming, not people fitting classifiers.
The mechanism: solver wrappers, implicit gradients, and vmap
The design separates the solver from the differentiation. A solver object exposes an update step that JAX can trace, and the library supplies an implicit-differentiation path so that the derivative of the solution with respect to problem inputs comes from the optimality conditions rather than from backpropagating through every iteration. The README describes this as differentiating "either implicitly or via autodiff of unrolled algorithm iterations," which is the real trade-off: unrolling is simple but memory grows with iteration count, while the implicit route needs the solver to converge and needs the optimality conditions to be well behaved.
Batchability comes from the same tracing machinery. Because problems are expressed as pure functions over arrays, vmap can lift a single-problem solver to a batch without rewriting it. Hardware acceleration is inherited from JAX and jaxlib rather than implemented by JAXopt itself, which is why the dependency floors in requirements.txt matter: jax>=0.2.18, jaxlib>=0.1.69, numpy>=1.18.4, scipy>=1.0.0. The package is pure Python, installed via setup.py with find_packages(), and the version string lives in jaxopt/version.py.
Installing jaxopt and running a first solver
The README gives three install commands. The released version comes from PyPI, the development version from the Git repository, and a source install uses setup.py. Whichever you pick, JAX itself should be installed first so that the right CPU, GPU or TPU build is selected.
$ pip install jaxoptFor the development version, the README gives this instead:
$ pip install git+https://github.com/google/jaxoptA source checkout installs with the command below. Note that the repository still ships setup.py rather than a pyproject.toml, so this is the path the project documents.
$ python setup.py installOnce installed, the entry points are the solver classes in the jaxopt package. The repository's examples/ directory is organized by problem type, with subdirectories for constrained, deep_learning, fixed_point and implicit_diff, plus examples/requirements.txt for the extra packages those notebooks and scripts need. Start from the example that matches your problem shape rather than from the API reference, because the interesting part of each example is how the solver is wrapped so that JAX can trace and differentiate it.
The maintenance status is the main limitation
The README's Status section is unambiguous: "JAXopt is no longer maintained nor developed." It points readers to the JAX website for alternatives and states that some features, specifically losses, projections and the LBFGS optimizer, have been ported into optax. The disclaimer adds that JAXopt was an open source project maintained by a dedicated team in Google Research and is not an official Google product.
The practical consequence is that the documentation is now the ceiling. If a solver fails to converge on your problem, there is no upstream fix coming. If a new JAX release changes tracing or transformation semantics, compatibility is not guaranteed by anyone. The last release listed is jaxopt-v0.8.5 from 2025-04-14, and the last push to the repository was on 2026-09-07, so the code is not frozen in a git sense, but the README's own statement is what you should plan around. Treat the library as a pinned dependency, not a moving one.
JAXopt vs optax, and why the split matters
The README names optax directly as where to look, and says losses, projections and the LBFGS optimizer have been ported there. That is a real difference in approach, not a rebranding. Optax is built around gradient transformations: you compose an optimizer from a chain of update rules and apply it to a loss. JAXopt is built around solving a problem, where the solver produces a solution and differentiation is attached to that solution.
If your work is training a model, optax is the intended home and the ported pieces cover the common cases. If your work is a constrained problem, a fixed point, or a bi-level objective where you need gradients through the solve, optax's gradient-transformation model does not express that, and JAXopt's implicit differentiation is the feature you are actually after. The comparison is therefore less "which is better" and more "which side of the solve do you need to differentiate."
Licence, dependencies and the cost of pinning
JAXopt is Apache-2.0, stated in the LICENSE file and in setup.py as license="Apache 2.0". That is a permissive licence, and the usual obligations apply: retain the licence and notices, and be aware of the patent grant and termination terms that Apache-2.0 carries. This is a description of the licence text, not legal advice; check it against your own distribution model.
The upgrade cost is the sharper issue. Because the project is unmaintained, every JAX upgrade is a compatibility question you answer yourself. The dependency floors in requirements.txt are old (jax>=0.2.18, jaxlib>=0.1.69), so they will not stop you installing a recent JAX, but they also will not warn you when a transformation JAXopt relies on changes. A workable approach is to pin jax, jaxlib and jaxopt together in the same environment and upgrade them as a set, testing the specific solvers you depend on rather than assuming the library tracks JAX.
Editorial conclusion
Adopt JAXopt only if you need implicit differentiation or a solver wrapper that optax does not provide, and only with the expectation that the repository will not receive fixes; the README states it is no longer maintained nor developed. Do not adopt it for plain neural-network training, where optax is the intended replacement. Before committing, check whether the specific solver you need (OSQP, LBFGS, a projection, a root finder) is already available in optax, and verify that your installed jax and jaxlib satisfy the floors in requirements.txt.
Frequently asked questions
Is JAXopt still maintained?
No. The README's Status section states that JAXopt is no longer maintained nor developed, and points readers to the JAX website for alternatives. Some features, including losses, projections and the LBFGS optimizer, have been ported into optax.
How do I install JAXopt?
The README gives three options: pip install jaxopt for the latest release, pip install git+https://github.com/google/jaxopt for the development version, or python setup.py install from a source checkout.
What is the difference between JAXopt and optax?
JAXopt is built around solving optimization problems and differentiating the solution, either implicitly or by unrolling iterations. Optax is built around gradient transformations for training. The README states that some JAXopt features have been ported into optax.
What problems can JAXopt solve?
The repository's examples directory is organized into constrained, deep_learning, fixed_point and implicit_diff, which reflects the problem types the library covers. The README describes the library as providing hardware accelerated, batchable and differentiable optimizers.
What are JAXopt's dependencies?
The requirements.txt file lists jax>=0.2.18, jaxlib>=0.1.69, numpy>=1.18.4 and scipy>=1.0.0. The package itself is distributed under the Apache 2.0 license.
Official sources
Add this badge to your README
If you maintain this project, the badge below links readers to this analysis and shows its maintenance status from the daily GitHub snapshot. Paste the markdown into your README; add ?metric=license or ?metric=stars to the image URL for a different field.
[](https://hysenlabs.com/projects/google-jaxopt)