NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #2402 most downloaded on PyPI
A gradient processing and optimization library in JAX.
Last release 6 months ago
20 Mar 2026
Ships fairly regularly
a new release about every 3 months
Some releases are documented
notes for 10 of 26 stable releases
Nothing withdrawn
no release was ever pulled
6 years old
26 releases · first in 2020
One column per quarter.
Following the JAX 0.9.2 release, the jax_pmap_shmap_merge config flag was removed so that the jax.pmap implementation is always based on jax.jit and j
jax_pmap_shmap_merge config flag was removed so that the jax.pmap implementation is always based on jax.jit and jax.shard_map, and opting into the old jax.pmap behavior is no longer an option. Optax had opted into the old behavior to give users time to migrate, and as of Optax 0.2.8 this is no longer supported. This changed shouldn't impact most users, but if you experience errors or performance regressions as a result of it, you can update your code following JAX's migration guide (or use JAX 0.9.2 or earlier and set jax.config.update("jax_pmap_shmap_merge", False)).Adversarial training example by @rajasekharporeddy in #1609inject_hyperparams uses the dtype inferred from parameters... by @copybara-service[bot] in #1615Full Changelog: v0.2.7...v0.2.8
Deprecate second order utilities. by @copybara-service [bot] in #1466
monitor and measure_with_ema to Optax transformations. by @copybara-service[bot] in #1430warn_deprecated_function instead of the Chex version. by @copybara-service[bot] in #1469assert_trees_all_{close, equal} functions and use them in tests instead of the Chex versions. by @copybara-service[bot] in #1481chex.Numeric with jax.typing.ArrayLike. Replace chex.Scalar with float/int as appropriate. by @copybara-service[bot] in #1479jax_pmap_shmap_merge=False flag in contrib/_complex_valued. #jax-fixit by @copybara-service[bot] in #1515ArrayTree and use it instead of chex.ArrayTree. by @copybara-service[bot] in #1484multiply_by_parameter_scale type in adafactor optimizer. by @copybara-service[bot] in #1504Lookahead Optimizer on MNIST example by @rajasekharporeddy in #1568ArrayTree and use it instead of chex.ArrayTree. by @copybara-service[bot] in #1588Full Changelog: v0.2.6...v0.2.7
absl.TestCase or parameterized.TestCase
instead of chex.TestCase as part of our effort to remove the chex
dependency. This means that Chex test variants (with/without jit, with/without
device_put, with pmap) are no longer tested. We decided it was sufficient to
use jit throughout the tests. There is already test coverage on both CPU and
accelerators, and pmap is deprecated.poly_loss_cross_entropy,
ctc_loss_with_forward_probs, ctc_loss, sigmoid_focal_loss) and regression
losses (huber_loss, cosine_similarity, cosine_distance) no longer support
positional args for hyperparameter-like inputs.chex types with internally-defined types and removed chex
dependency.Fix for #1328 by @copybara-service [bot] in #1329
Freezing in transformations api page in the documentation by @rajasekharporeddy in #1331optax.transforms in freezing documentation by @rajasekharporeddy in #1334Full Changelog: v0.2.5...v0.2.6
Remove longtime deprecated functions. by @copybara-service in #1149
optax.partition. by @copybara-service in #1150contrib by @leloykun in #1126scalar_type_of - random.key breaks it by @rdyro in #1169contrib by @mathDR in #1104freeze, selective_optimizer, and Prefix-Mask Support by @pranavagrawaI in #1284freeze and selective_transform by @pranavagrawaI in #1300Note truncated.
Replace deprecated typing.Hashable with collections.abc.Hashable. by @carlosgmartin in #1068
multi_normal from utilities.rst by @Abhinavcode13 in #1009optax/monte_carlo. by @copybara-service in #1076Full Changelog: v0.2.3...v0.2.4
Fix jax.tree_map deprecation warnings. by @copybara-service in #917
jax.tree_* functions with jax.tree.* by @copybara-service in #963normalize_by_update_norm by @SauravMaheshkar in #958Full Changelog: v0.2.2...v0.2.3
Added mathematical description to Noisy SGD by @hmludwig in #857
reduce_on_plateau.ipynb:20002: WARNING: No source code lexer found for notebook cell by @copybara-service in #875Full Changelog: v0.2.1...v0.2.2
Begin of 0.2.1 development by @copybara-service in #845
Full Changelog: v0.2.0...v0.2.1
Begin of 0.2.0 development by @copybara-service in #766
adversarial_training.ipynb by @copybara-service in #776params by @yixiaoer in #785cosine_decay_schedule. by @copybara-service in #797Full Changelog: v0.1.9...v0.2.0
Deprecate optax.inject_stateful_hyperparameters and replace it with optax.inject_hyperparameters by @copybara-service in #730
atol option to contrib.reduce_on_plateau() by @stefanocortinovis in #698versionadded and seealso metadata to (n)adam(w) solvers by @copybara-service in #729setup.py from publishing CI by @SauravMaheshkar in #737Full Changelog: v0.1.8...v0.1.9
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Nothing published for this version
Your coding agent can read these notes before it upgrades. Set up the MCP server →