$npx skillfedfor your agent

optax

A gradient processing and optimization library in JAX.

Worth itPyPI Python ModulesReleased Mar 20263.8M downloads / mopermissive licensePure Python

Decision gist · record as of 2026-08-14

pure-Python wheel — optax-0.2.8-py3-none-any.whl
v0.2.8 · released 2026-03-20 · Python >=3.10 · 4 runtime deps: absl-py, jax, jaxlib, numpy

Yes. Optax is actively maintained, has no known vulnerabilities, installs with low friction, and is the standard gradient optimization library for JAX. Install it if you are building machine learning systems with JAX and need flexible, composable optimizer and loss components. The only gotcha is the compiled JAX/jaxlib dependency, which may require platform-specific setup.AI-flagged interpretation of the facts on this page — verify before relying

Before you install

  • Requires JAX and jaxlib (compiled dependency); Python >= 3.10.
  • Low friction install as a pure Python wheel.
  • Actively maintained with recent releases; last commit 2026-08-07 and latest release 2026-03-20.

License · maintenance · safety

permissive license (permissive) — Licensed under Apache Software License (permissive). No restrictions on commercial or private use.

last release 2026-03-20 (147 days) · last repo commit 2026-08-07 · 2,317 stars

0 known vulnerabilities (OSV.dev, 2026-08-14) · 3,760,118 downloads/mo, #2,503 on PyPI

Verify before relying

pip install optax

import optax
import jax
import jax.numpy as jnp

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)
  • Whether the library's optimizer implementations have been benchmarked against standard baselines in recent comparisons.
  • Performance characteristics when used with large-scale models or distributed training setups.
Same gist for agents: .md · .json

What it is and what it does

Optax is a gradient processing and optimization library built on top of JAX. It provides well-tested, efficient implementations of core optimization components—such as Adam, SGD, and other popular optimizers—along with loss functions and gradient transformation utilities. The library is designed around composability: rather than monolithic optimizer classes, it offers small building blocks that can be recombined in custom ways to create new optimizers or gradient processing pipelines.

The package targets researchers and practitioners working with JAX who need flexible, modular optimization tools. It depends on JAX, jaxlib, numpy, and absl-py. The library evolved from an earlier experimental JAX module and is now maintained as a standalone project by DeepMind. It supports current Python versions and is actively maintained, making it suitable for both research prototyping and production use in JAX-based machine learning workflows.

Use it for

  • Training neural networks with custom optimizer combinations by composing Optax building blocks.
  • Implementing gradient clipping, weight decay, or other transformations in a modular way.
  • Prototyping new optimization algorithms by combining existing components.
  • Using standard optimizers like Adam or RMSprop in JAX-based machine learning projects.
  • Computing loss functions like L2 or cross-entropy within JAX training loops.

Worth the install?

AI-flagged interpretation of the facts on this page. Verify before relying on it.

Worth it

Yes.

Optax is actively maintained, has no known vulnerabilities, installs with low friction, and is the standard gradient optimization library for JAX. Install it if you are building machine learning systems with JAX and need flexible, composable optimizer and loss components. The only gotcha is the compiled JAX/jaxlib dependency, which may require platform-specific setup.

Install

optax on PyPI

Before you install

Low friction install as a pure Python wheel. Actively maintained with recent releases; last commit 2026-08-07 and latest release 2026-03-20. Requires JAX and its compiled dependency jaxlib, which may add setup complexity depending on your platform.

Requires JAX and jaxlib (compiled dependency); Python >= 3.10.

License in practice

Licensed under Apache Software License (permissive). No restrictions on commercial or private use.

Quickstart

pip install optax

import optax
import jax
import jax.numpy as jnp

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)

Verify before relying

  • Whether the library's optimizer implementations have been benchmarked against standard baselines in recent comparisons.
  • Performance characteristics when used with large-scale models or distributed training setups.

Package facts

Licensepermissive license permissive
Python supportSupports the current Python release >=3.10
Install frictionLow. Pure-Python wheel
Runtime dependencies
4 packages
absl-pyjaxjaxlibnumpy
MaintenanceActively maintained 147 days since the last release
Last repo commit
First released
Downloads3,760,118 / month, #2,503 on PyPI 30-day window, as of 2026-08-14
Known vulnerabilitiesNone known OSV.dev, checked 2026-08-14
Classifiers
Development Status :: 4 - BetaEnvironment :: ConsoleIntended Audience :: DevelopersIntended Audience :: Science/ResearchLicense :: OSI Approved :: Apache Software LicenseOperating System :: OS IndependentProgramming Language :: PythonProgramming Language :: Python :: 3Topic :: Scientific/Engineering :: Artificial IntelligenceTopic :: Software Development :: Libraries :: Python Modules

Evidence: optax-0.2.8-py3-none-any.whl

Tags

Capabilities
JAX optimizer librarygradient processing JAXAdam optimizer JAXcustom optimizers JAXmachine learning optimizationneural network training JAXloss functions JAX
Topics
jax-ecosystemgradient-optimizationcomposable-components
PyPI keywords
pythonmachine learningreinforcement-learning

Let your AI agent find packages like this

Example. Real query, live index.

You found this page by searching. An agent finds it by wishing: SkillFed indexes 14,416 PyPI packages by what they can do, searchable in plain language.

wish › “JAX optimizer library”

  • optaxOptax provides composable building blocks for gradient processing and…
  • orbax-checkpointOrbax Checkpoint provides asynchronous checkpointing for JAX machine…
  • saxSAX is a JAX-based simulator and optimizer for S-parameter circuits,…

Give your agent the search over MCP, or paste the wish link into any chat.

More Python Modules packages

idna Worth it
PyPI · Python Modules · released Jun 2026

Converts domain names between Unicode and ASCII-compatible encoding (Punycode) according to IDNA 2008 and Unicode Technical Standard 46, with security validation and broader script coverage than the standard library.

Install it if you work with internationalized domain names, need to validate domains, or use HTTP clients that depend on it transitively.

BSD-3-Clausepure Python · 3.9+
1.8Bdownloads / mo
setuptools Worth it
PyPI · Python Modules · released Aug 2026

Setuptools is a Python build backend and package management tool that handles building, distributing, and installing Python packages, including support for C/C++ extension modules.

MITpure Python · 3.10+
1.6Bdownloads / mo
PyYAML Worth it
PyPI · Python Modules · released Sep 2025

PyYAML parses and emits YAML 1.1 data format, enabling serialization and deserialization of configuration files and Python objects to and from human-readable YAML text.

MITcompiled wheel · 3.8+
1.2Bdownloads / mo
pydantic Worth it
PyPI · Python Modules · released May 2026

Pydantic validates Python data structures against type hints, coercing and checking input at runtime to ensure it matches a declared schema.

MITpure Python · 3.9+
1.1Bdownloads / mo
annotated-types Worth it
PyPI · Python Modules · released Jul 2026

Provides reusable metadata objects for use with PEP-593 `typing.Annotated` to express common constraints like bounds, collection sizes, and predicates on types.

Install it if you use or build libraries that need to express type constraints in a standardized, inspectable way—or if you want to annotate your own types with…

MITpure Python · 3.10+
871.3Mdownloads / mo
typing-inspection Worth it
PyPI · Python Modules · released Aug 2026

Provides runtime tools to inspect and introspect Python type annotations, enabling programmatic examination of type hints at execution time.

MITpure Python · 3.10+
783.0Mdownloads / mo

See also jmp · optimistix · equinox · lineax · diffrax · prodigyopt · jax · pytorch_optimizer · ropt · pyswarms

Further reading