NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #1018 most downloaded on PyPI
Differentiate, compile, and transform Numpy code.
Last release today
17 Sep 2026
Ships on a steady schedule
a new release about every 4 weeks
Nearly every release is documented
notes for 60 of the last 60 stable releases
5 versions withdrawn
withdrawn after publishing
8 years old
196 releases · first in 2018
Deleted jax.experimental.callback
Breaking changes
jax.experimental.callbacknp.ndarray now can raise errors when the result is used as a shape value
({jax-issue}#14106).14102)Changes
jax2tf.call_tf has a new parameter has_side_effects (default True)
that can be used to declare whether an instance can be removed or replicated
by JAX optimizations such as dead-code elimination ({jax-issue}#13980).#14108).jax.Array has been enabled by default in JAX 0.4 and makes some breaking change to the pjit API. The jax.Array migration guide can help you migrate yo…
One column per quarter.
version-support-policy.jax.Array which is a unified array type that subsumes
DeviceArray, ShardedDeviceArray, and GlobalDeviceArray types in JAX.
The jax.Array type helps make parallelism a core feature of JAX,
simplifies and unifies JAX internals, and allows us to unify jit and
pjit. jax.Array has been enabled by default in JAX 0.4 and makes some
breaking change to the pjit API. The jax.Array migration
guide can
help you migrate your codebase to jax.Array. You can also look at the
Distributed arrays and automatic parallelization
tutorial to understand the new concepts.PartitionSpec and Mesh are now out of experimental. The new API endpoints
are jax.sharding.PartitionSpec and jax.sharding.Mesh.
jax.experimental.maps.Mesh and jax.experimental.PartitionSpec are
deprecated and will be removed in 3 months.with_sharding_constraints new public endpoint is
jax.lax.with_sharding_constraint.jax.config, the ABSL flag values are no
longer read or written after the JAX configuration options are initially
populated from the ABSL flags. This change improves performance of reading
jax.config options, which are used pervasively in JAX.jax.numpy functions now have their arguments marked as
positional-only, matching NumPy.jnp.msort is now deprecated, following the deprecation of np.msort in numpy 1.24.
It will be removed in a future release, in accordance with the {ref}api-compatibility
policy. It can be replaced with jnp.sort(a, axis=0).* The release was yanked.
{func}jax.numpy.linalg.pinv now supports the hermitian option.
jax.numpy.linalg.pinv now supports the hermitian option.jax.scipy.linalg.hessenberg is now supported on CPU only. Requires
jaxlib > 0.3.24.jax.lax.linalg.hessenberg,
{func}jax.lax.linalg.tridiagonal, and
{func}jax.lax.linalg.householder_product were added. Householder reduction
is currently CPU-only and tridiagonal reductions are supported on CPU and
GPU only.svd and jax.numpy.linalg.pinv are now computed more
economically for non-square matrices.jax_experimental_name_stack config option.axis_names arguments to the
{class}jax.experimental.maps.Mesh constructor into a singleton tuple
instead of unpacking the string into a sequence of character axis names.JAX should be faster to import. We now import scipy lazily, which accounted for a significant fraction of JAX's import time.
JAX_PERSISTENT_CACHE_MIN_COMPILE_TIME_SECS=$N can be
used to limit the number of cache entries written to the persistent cache.
By default, computations that take 1 second or more to compile will be
cached.
jax.scipy.stats.mode.pmap on TPU if no order is specified now
matches jax.devices() for single-process jobs. Previously the
two orderings differed, which could lead to unnecessary copies or
out-of-memory errors. Requiring the orderings to agree simplifies matters.jax.numpy.gradient now behaves like most other functions in {mod}jax.numpy,
and forbids passing lists or tuples in place of arrays ({jax-issue}#12958)jax.numpy.linalg and {mod}jax.numpy.fft now uniformly
require inputs to be array-like: i.e. lists and tuples cannot be used in place
of arrays. Part of {jax-issue}#7737.jax.sharding.MeshPspecSharding has been renamed to jax.sharding.NamedSharding.
jax.sharding.MeshPspecSharding name will be removed in 3 months.Update Colab TPU driver version for new jaxlib release.
Several test utilities deprecated in JAX v0.3.8 are now removed from {mod}jax.test_util.
JAX_PLATFORMS=tpu,cpu as default setting in TPU initialization,
so JAX will raise an error if TPU cannot be initialized instead of falling
back to CPU. Set JAX_PLATFORMS='' to override this behavior and automatically
choose an available backend (the original default), or set JAX_PLATFORMS=cpu
to always use CPU regardless of if the TPU is available.jax.test_util.The persistent compilation cache will now warn instead of raising an exception on error ({jax-issue}#12582), so program execution can continue if some
#12582), so program execution can continue
if something goes wrong with the cache. Set
JAX_RAISE_PERSISTENT_CACHE_ERRORS=true to revert this behavior.GitHub commits .
Changes
The persistent compilation cache will now warn instead of raising an exception on error ( #12582 ), so program execution can continue if something goes wrong with the cache. Set JAX_RAISE_PERSISTENT_CACHE_ERRORS=true to revert this behavior.
Adds missing .pyi files that were missing from the previous release (#12536).
Notable changes:
.pyi files that were missing from the previous release (#12536).jax 0.3.19 and the libtpu version it pinned (#12550). Requires jaxlib 0.3.20.pip url in setup.py comment (#12528)..pyi files that were missing from the previous release ({jax-issue}#12536).jax 0.3.19 and the libtpu version it pinned ({jax-issue}#12550). Requires jaxlib 0.3.20.pip url in setup.py comment ({jax-issue}#12528).Fixes the required jaxlib version
Fixes the required jaxlib version
jax.soft_pmap has been deleted. Please use pjit or xmap instead. jax.soft_pmap is undocumented. If it were documented, a deprecation period would have…
#7733) is stable and public. See the
overview and the API docs
for {mod}jax.stages.jax.Array, intended to be used for both isinstance checks
and type annotations for array types in JAX. Notice that this included some subtle
changes to how isinstance works for {class}jax.numpy.ndarray for jax-internal
objects, as {class}jax.numpy.ndarray is now a simple alias of {class}jax.Array.jax._src is no longer imported into the from the public jax namespace.
This may break users that were using JAX internals.jax.soft_pmap has been deleted. Please use pjit or xmap instead.
jax.soft_pmap is undocumented. If it were documented, a deprecation period
would have been provided.GitHub commits .
Changes
Ahead-of-time lowering and compilation functionality (tracked in #7733 ) is stable and public. See the overview and the API docs for jax.stages .
Introduced jax.Array , intended to be used for both isinstance checks and type annotations for array types in JAX. Notice that this included some subtle changes to how isinstance works for jax.numpy.ndarray for jax-internal objects, as jax.numpy.ndarray is now a simple alias of jax.Array .
Breaking changes
jax._src is no longer imported into the public jax namespace. This may break users that were using JAX internals.
jax.soft_pmap has been deleted. Please use pjit or xmap instead. jax.soft_pmap is undocumented. If it were documented, a deprecation period would have been provided.
jax.checkpoint, also known as jax.remat, no longer supports the concrete option, following the previous version's deprecation; see JEP 11830.
lax.pow with an exponent of zero (#12041)jax.checkpoint, also known as jax.remat, no longer supports the concrete option, following the previous version's deprecation; see JEP 11830.jax.pure_callback that enables calling back to pure Python functions from compiled functions (e.g. functions decorated with jax.jit or jax.pmap).DeviceArray.tile() method has been removed. Use jax.numpy.tile (#11944).DeviceArray.to_py() has been deprecated. Use np.asarray(x) instead.lax.pow with an exponent of zero
({jax-issue}12041)jax.checkpoint, also known as {func}jax.remat, no longer supports
the concrete option, following the previous version's deprecation; see
JEP 11830.jax.pure_callback that enables calling back to pure Python functions from compiled functions (e.g. functions decorated with jax.jit or jax.pmap).DeviceArray.tile() method has been removed. Use {func}jax.numpy.tile
({jax-issue}#11944).DeviceArray.to_py() has been deprecated. Use np.asarray(x) instead.Support for NumPy 1.19 has been dropped, per the deprecation policy. Please upgrade to NumPy 1.20 or newer.
jax.debug that includes utilities for runtime value debugging such at jax.debug.print and jax.debug.breakpoint.jax.mask jax.shapecheck APIs have been removed. See #11557.jax.experimental.loops has been removed. See #10278 for an alternative API.jax.tree_util.tree_multimap has been removed. It has been deprecated since JAX release 0.3.5, and jax.tree_util.tree_map is a direct replacement.jax.experimental.stax; it has long been a deprecated alias of jax.example_libraries.stax.jax.experimental.optimizers; it has long been a deprecated alias of jax.example_libraries.optimizers.jax.checkpoint, also known as jax.remat, has a new implementation switched on by default, meaning the old implementation is deprecated; see JEP 11830.jax.debug that includes utilities for runtime value debugging such at {func}jax.debug.print and {func}jax.debug.breakpoint.jax.mask {func}jax.shapecheck APIs have been removed.
See {jax-issue}#11557.jax.experimental.loops has been removed. See {jax-issue}#10278
for an alternative API.jax.tree_util.tree_multimap has been removed. It has been deprecated since
JAX release 0.3.5, and {func}jax.tree_util.tree_map is a direct replacement.jax.experimental.stax; it has long been a deprecated alias of
{mod}jax.example_libraries.stax.jax.experimental.optimizers; it has long been a deprecated alias of
{mod}jax.example_libraries.optimizers.jax.checkpoint, also known as {func}jax.remat, has a new
implementation switched on by default, meaning the old implementation is
deprecated; see JEP 11830.JaxTestCase and JaxTestLoader have been removed from jax.test_util. These classes have been deprecated since v0.3.1 ({jax-issue}#11248).
JaxTestCase and JaxTestLoader have been removed from jax.test_util. These
classes have been deprecated since v0.3.1 ({jax-issue}#11248).jax.scipy.gaussian_kde ({jax-issue}#11237).dict, list, set, tuple)
now raise a TypeError in all cases. Previously some cases (particularly equality and inequality)
would return boolean scalars inconsistent with similar operations in NumPy ({jax-issue}#11234).jax.tree_util routines accessed as top-level JAX package imports are now
deprecated, and will be removed in a future JAX release in accordance with the
{ref}api-compatibility policy:
jax.treedef_is_leaf is deprecated in favor of {func}jax.tree_util.treedef_is_leafjax.tree_flatten is deprecated in favor of {func}jax.tree_util.tree_flattenjax.tree_leaves is deprecated in favor of {func}jax.tree_util.tree_leavesjax.tree_structure is deprecated in favor of {func}jax.tree_util.tree_structurejax.tree_transpose is deprecated in favor of {func}jax.tree_util.tree_transposejax.tree_unflatten is deprecated in favor of {func}jax.tree_util.tree_unflattensym_pos argument of {func}jax.scipy.linalg.solve is deprecated in favor of assume_a='pos',
following a similar deprecation in {func}scipy.linalg.solve.GitHub commits .
Changes
JaxTestCase and JaxTestLoader have been removed from jax.test_util . These classes have been deprecated since v0.3.1 ( #11248 ).
Added jax.scipy.gaussian_kde ( #11237 ).
Binary operations between JAX arrays and built-in collections ( dict , list , set , tuple ) now raise a TypeError in all cases. Previously some cases (particularly equality and inequality) would return boolean scalars inconsistent with similar operations in NumPy ( #11234 ).
Several jax.tree_util routines accessed as top-level JAX package imports are now deprecated, and will be removed in a future JAX release in accordance with the API compatibility policy:
jax.treedef_is_leaf() is deprecated in favor of jax.tree_util.treedef_is_leaf()
jax.tree_flatten() is deprecated in favor of jax.tree_util.tree_flatten()
jax.tree_leaves() is deprecated in favor of jax.tree_util.tree_leaves()
jax.tree_structure() is deprecated in favor of jax.tree_util.tree_structure()
jax.tree_transpose() is deprecated in favor of jax.tree_util.tree_transpose()
jax.tree_unflatten() is deprecated in favor of jax.tree_util.tree_unflatten()
The sym_pos argument of jax.scipy.linalg.solve() is deprecated in favor of assume_a='pos' , following a similar deprecation in scipy.linalg.solve() .
In scatter-update operations (i.e. {attr}jax.numpy.ndarray.at), unsafe implicit dtype casts are deprecated, and now result in a FutureWarning. In a fu…
jax.experimental.compilation_cache.initialize_cache does not support
max_cache_size_ bytes anymore and will not get that as an input.JAX_PLATFORMS now raises an exception when platform initialization fails.jax.numpy.linalg.slogdet now accepts an optional method argument
that allows selection between an LU-decomposition based implementation and
an implementation based on QR decomposition.jax.numpy.linalg.qr now supports mode="raw".pickle, copy.copy, and copy.deepcopy now have more complete support when
used on jax arrays ({jax-issue}#10659). In particular:
pickle and deepcopy previously returned np.ndarray objects when used
on a DeviceArray; now DeviceArray objects are returned. For deepcopy,
the copied array is on the same device as the original. For pickle the
deserialized array will be on the default device.deepcopy and copy
previously were no-ops. Now they use the same mechanism as DeviceArray.copy().pickle on a traced array now results in an explicit
ConcretizationTypeError.jax.numpy.ldexp no longer silently promotes all inputs to float64,
instead it promotes to float32 for integer inputs of size int32 or smaller
({jax-issue}#10921).create_perfetto_link option to {func}jax.profiler.start_trace and
{func}jax.profiler.start_trace. When used, the profiler will generate a
link to the Perfetto UI to view the trace.jax.profiler.start_server(...) to store the
keepalive globally, rather than requiring the user to keep a reference to
it.jax.random.generalized_normal.jax.random.ball.jax.default_device.python -m jax.collect_profile script to manually capture program
traces as an alternative to the TensorBoard UI.jax.named_scope context manager that adds profiler metadata to
Python programs (similar to jax.named_call).jax.numpy.ndarray.at), unsafe implicit
dtype casts are deprecated, and now result in a FutureWarning.
In a future release, this will become an error. An example of an unsafe implicit
cast is jnp.zeros(4, dtype=int).at[0].set(1.5), in which 1.5 previously was
silently truncated to 1.jax.experimental.compilation_cache.initialize_cache now supports gcs
bucket path as input.jax.scipy.stats.gennorm.jax.numpy.roots is now better behaved when strip_zeros=False when
coefficients have leading zeros ({jax-issue}#11215).* GitHub commits.
Fixes https://github.com/google/jax/pull/10717
{func}jax.scipy.linalg.polar_unitary, which was a JAX extension to the scipy API, is deprecated. Use {func}jax.scipy.linalg.polar instead.
jax.lax.eigh now accepts an optional sort_eigenvalues argument
that allows users to opt out of eigenvalue sorting on TPU.jax.lax.linalg are now marked
keyword-only. As a backward-compatibility step passing keyword-only
arguments positionally yields a warning, but in a future JAX release passing
keyword-only arguments positionally will fail.
However, most users should prefer to use {mod}jax.numpy.linalg instead.jax.scipy.linalg.polar_unitary, which was a JAX extension to the
scipy API, is deprecated. Use {func}jax.scipy.linalg.polar instead.GitHub commits .
Changes
jax.lax.eigh() now accepts an optional sort_eigenvalues argument that allows users to opt out of eigenvalue sorting on TPU.
Deprecations
Non-array arguments to functions in jax.lax.linalg are now marked keyword-only. As a backward-compatibility step passing keyword-only arguments positionally yields a warning, but in a future JAX release passing keyword-only arguments positionally will fail. However, most users should prefer to use jax.numpy.linalg instead.
jax.scipy.linalg.polar_unitary() , which was a JAX extension to the scipy API, is deprecated. Use jax.scipy.linalg.polar() instead.
* GitHub commits.
Added support for fully asynchronous checkpointing for GlobalDeviceArray.
Many functions and objects available in {mod}jax.test_util are now deprecated and will raise a warning on import. This includes cases_from_list, check…
jax.numpy.linalg.svd on TPUs uses a qdwh-svd solver.jax.numpy.linalg.cond on TPUs now accepts complex input.jax.numpy.linalg.pinv on TPUs now accepts complex input.jax.numpy.linalg.matrix_rank on TPUs now accepts complex input.jax.scipy.cluster.vq.vq has been added.jax.experimental.maps.mesh has been deleted.
Please use jax.experimental.maps.Mesh. Please see https://jax.readthedocs.io/en/latest/_autosummary/jax.experimental.maps.Mesh.html#jax.experimental.maps.Mesh
for more information.jax.scipy.linalg.qr now returns a length-1 tuple rather than the raw array when mode='r', in order to match the behavior of scipy.linalg.qr ({jax-issue}#10452)jax.numpy.take_along_axis now takes an optional mode parameter that specifies the behavior of out-of-bounds indexing. By default, invalid values (e.g., NaN) will be returned for out-of-bounds indices. In previous versions of JAX, invalid indices were clamped into range. The previous behavior can be restored by passing mode="clip".jax.numpy.take now defaults to mode="fill", which returns invalid values (e.g., NaN) for out-of-bounds indices.x.at[...].set(...), now have "drop" semantics. This has no effect on the scatter operation itself, but it means that when differentiated the gradient of a scatter will yield zero cotangents for out-of-bounds indices. Previously out-of-bounds indices were clamped into range for the gradient, which was not mathematically correct.jax.numpy.take_along_axis now raises a TypeError if its indices are not of an integer type, matching the behavior of
{func}numpy.take_along_axis. Previously non-integer indices were silently cast to integers.jax.numpy.ravel_multi_index now raises a TypeError if its dims argument is not of an integer type, matching the behavior of {func}numpy.ravel_multi_index. Previously non-integer dims was silently cast to integers.jax.numpy.split now raises a TypeError if its axis argument is not of an integer type, matching the behavior of {func}numpy.split. Previously non-integer axis was silently cast to integers.jax.numpy.indices now raises a TypeError if its dimensions are not of an integer type, matching the behavior of {func}numpy.indices. Previously non-integer dimensions were silently cast to integers.jax.numpy.diag now raises a TypeError if its k argument is not of an integer type, matching the behavior of {func}numpy.diag. Previously non-integer k was silently cast to integers.jax.random.orthogonal.jax.test_util are now deprecated and will raise a warning on import. This includes cases_from_list, check_close, check_eq, device_under_test, format_shape_dtype_string, rand_uniform, skip_on_devices, with_config, xla_bridge, and _default_tolerance ({jax-issue}#10389). These, along with previously-deprecated JaxTestCase, JaxTestLoader, and BufferDonationTestCase, will be removed in a future JAX release. Most of these utilites can be replaced by calls to standard python & numpy testing utilities found in e.g. {mod}unittest, {mod}absl.testing, {mod}numpy.testing, etc. JAX-specific functionality such as device checking can be replaced through the use of public APIs such as {func}jax.devices. Many of the deprecated utilities will still exist in {mod}jax._src.test_util, but these are not public APIs and as such may be changed or removed without notice in future releases.The DeviceArray.tile() method is deprecated, because numpy arrays do not have a tile() method. As a replacement for this, use jax.numpy.tile (#10266).
jax.numpy.take_along_axis were broadcasted (#10281).jax.scipy.special.expit and jax.scipy.special.logit now require their arguments to be scalars or JAX arrays. They also now promote integer arguments to floating point.DeviceArray.tile() method is deprecated, because numpy arrays do not have a tile() method. As a replacement for this, use jax.numpy.tile (#10266).jax.numpy.take_along_axis were broadcasted ({jax-issue}#10281).jax.scipy.special.expit and {func}jax.scipy.special.logit now
require their arguments to be scalars or JAX arrays. They also now promote
integer arguments to floating point.DeviceArray.tile() method is deprecated, because numpy arrays do not have a
tile() method. As a replacement for this, use {func}jax.numpy.tile
({jax-issue}#10266).Upgraded libtpu wheel to the fixed version. Fixes #10218.
jax.experimental.loops is being deprecated. See {jax-issue}#10278
for an alternative API.GitHub commits .
Changes:
Upgraded libtpu wheel to a version that fixes a hang when initializing a TPU pod. Fixes #10218 .
Deprecations:
jax.experimental.loops is being deprecated. See #10278 for an alternative API.
jax.nn.normalize is being deprecated. Use jax.nn.standardize instead (#9899).
Changes
jax.random.loggamma & improved behavior of jax.random.beta
and jax.random.dirichlet for small parameter values (#9906).lax_numpy submodule is no longer exposed in the jax.numpy namespace (#10029).jax.numpy.frombuffer, jax.numpy.fromfunction,
and jax.numpy.fromstring (#10049).DeviceArray.copy() now returns a DeviceArray rather than a np.ndarray (#10069)jax.scipy.linalg.rsf2csfjax.nn.normalize is being deprecated. Use jax.nn.standardize instead (#9899).jax.tree_util.tree_multimap is deprecated. Use jax.tree_util.tree_map instead (#5746).jax.experimental.sharded_jit is deprecated. Use pjit instead.jax.random.loggamma & improved behavior of {func}jax.random.beta
and {func}jax.random.dirichlet for small parameter values ({jax-issue}#9906).lax_numpy submodule is no longer exposed in the jax.numpy namespace ({jax-issue}#10029).jax.numpy.frombuffer, {func}jax.numpy.fromfunction,
and {func}jax.numpy.fromstring ({jax-issue}#10049).DeviceArray.copy() now returns a DeviceArray rather than a np.ndarray ({jax-issue}#10069)jax.scipy.linalg.rsf2csfjax.experimental.sharded_jit has been deprecated and will be removed soon.jax.nn.normalize is being deprecated. Use {func}jax.nn.standardize instead ({jax-issue}#9899).jax.tree_util.tree_multimap is deprecated. Use {func}jax.tree_util.tree_map instead ({jax-issue}#5746).jax.experimental.sharded_jit is deprecated. Use pjit instead.GitHub commits .
Changes:
added jax.random.loggamma() & improved behavior of jax.random.beta() and jax.random.dirichlet() for small parameter values ( #9906 ).
the private lax_numpy submodule is no longer exposed in the jax.numpy namespace ( #10029 ).
added array creation routines jax.numpy.frombuffer() , jax.numpy.fromfunction() , and jax.numpy.fromstring() ( #10049 ).
DeviceArray.copy() now returns a DeviceArray rather than a np.ndarray ( #10069 )
added jax.scipy.linalg.rsf2csf()
jax.experimental.sharded_jit has been deprecated and will be removed soon.
Deprecations:
jax.nn.normalize() is being deprecated. Use jax.nn.standardize() instead ( #9899 ).
jax.tree_util.tree_multimap() is deprecated. Use jax.tree_util.tree_map() instead ( #5746 ).
jax.experimental.sharded_jit is deprecated. Use pjit instead.
Fix a bug introduced in #9923.
Fix a bug introduced in #9923.
* GitHub commits.
The functions jax.ops.index_update, jax.ops.index_add, which were deprecated in 0.2.22, have been removed. Please use the `.at` property on JAX arrays…
jax.ops.index_update, jax.ops.index_add, which were
deprecated in 0.2.22, have been removed. Please use
the .at property on JAX arrays
instead, e.g., x.at[idx].set(y).jax.experimental.ann.approx_*_k into jax.lax. These functions are
optimized alternatives to jax.lax.top_k.jax.numpy.broadcast_arrays and {func}jax.numpy.broadcast_to now require scalar
or array-like inputs, and will fail if they are passed lists (part of {jax-issue}#7737).pjit now works on CPU (in addition to previous TPU and GPU support).jax.test_util.JaxTestCase and jax.test_util.JaxTestLoader are now deprecated. The suggested replacement is to use parametrized.TestCase directly. For…
jax.test_util.JaxTestCase and jax.test_util.JaxTestLoader are now deprecated.
The suggested replacement is to use parametrized.TestCase directly. For tests that
rely on custom asserts such as JaxTestCase.assertAllClose(), the suggested replacement
is to use standard numpy testing utilities such as numpy.testing.assert_allclose(),
which work directly with JAX arrays (#9620 ).jax.test_util.JaxTestCase now sets jax_numpy_rank_promotion='raise' by default
(#9562 ). To recover the previous behavior, use the new
jax.test_util.with_config decorator:@jtu.with_config(jax_numpy_rank_promotion='allow')
class MyTestCase(jtu.JaxTestCase):
...
jax.scipy.linalg.schur, jax.scipy.linalg.sqrtm,
jax.scipy.signal.csd, jax.scipy.signal.stft,
jax.scipy.signal.welch.Changes:
jax.test_util.JaxTestCase and jax.test_util.JaxTestLoader are now deprecated.
The suggested replacement is to use parametrized.TestCase directly. For tests that
rely on custom asserts such as JaxTestCase.assertAllClose(), the suggested replacement
is to use standard numpy testing utilities such as {func}numpy.testing.assert_allclose(),
which work directly with JAX arrays ({jax-issue}#9620).jax.test_util.JaxTestCase now sets jax_numpy_rank_promotion='raise' by default
({jax-issue}#9562). To recover the previous behavior, use the new
jax.test_util.with_config decorator:@jtu.with_config(jax_numpy_rank_promotion='allow')
class MyTestCase(jtu.JaxTestCase):
...
jax.scipy.linalg.schur, {func}jax.scipy.linalg.sqrtm,
{func}jax.scipy.signal.csd, {func}jax.scipy.signal.stft,
{func}jax.scipy.signal.welch.jax version has been bumped to 0.3.0. Please see the design doc for the explanation.
Changes
jax.jit(f).lower(...).compiler_ir() now defaults to the MHLO dialect if no dialect= is passed.
jax.jit(f).lower(...).compiler_ir() now defaults to the MHLO dialect if no
dialect= is passed.jax.jit(f).lower(...).compiler_ir(dialect='mhlo') now returns an MLIR
ir.Module object instead of its string representation.Support for NumPy 1.18 has been dropped, per the deprecation policy. Please upgrade to a supported NumPy version.
Breaking changes:
JAX_HOST_CALLBACK_AD_TRANSFORMS environment variable, or the --flax_host_callback_ad_transforms flag. Additionally, added documentation for how to implement the old behavior using JAX custom AD APIs ({jax-issue}#8678).0.0 and NaN regardless of the bit representation. In particular, 0.0 and -0.0 are now treated as equivalent, where previously -0.0 was treated as less than 0.0. Additionally all NaN representations are now treated as equivalent and sorted to the end of the array. Previously negative NaN values were sorted to the front of the array, and NaN values with different internal bit representations were not treated as equivalent, and were sorted according to those bit patterns ({jax- issue}#9178).jax.numpy.unique now treats NaN values in the same way as np.unique in NumPy versions 1.21 and newer: at most one NaN value will appear in the uniquified output ({jax-issue}9184).Bug fixes:
#8907).New features:
jax.block_until_ready ({jax-issue}`#8941)JAX_DUMP_IR_TO=/path. If set, JAX dumps the MHLO/HLO IR it generates for each computation to a file under the given path.jax.ensure_compile_time_eval to the public api ({jax-issue}#7987).#9189).Out-of-bounds indices to jax.ops.segment_sum will now be handled with FILL_OR_DROP semantics, as documented. This primarily afects the reverse-mode de
Bug fixes:
Out-of-bounds indices to jax.ops.segment_sum will now be handled with FILL_OR_DROP semantics, as documented. This primarily afects the reverse-mode derivative, where gradients corresponding to out-of-bounds indices will now be returned as 0. (#8634).
jax2tf will force the converted code to use XLA for the code fragments under jax.jit, e.g., most jax.numpy functions (#7839).
Bug fixes:
jax.ops.segment_sum will now be handled with
FILL_OR_DROP semantics, as documented. This primarily affects the
reverse-mode derivative, where gradients corresponding to out-of-bounds
indices will now be returned as 0. (#8634).#7839).(Experimental) jax.distributed.initialize exposes multi-host GPU backend.
New features:
jax.distributed.initialize exposes multi-host GPU backend.jax.random.permutation supports new independent keyword argument
({jax-issue}#8430)Breaking changes
jax.experimental.stax to jax.example_libraries.staxjax.experimental.optimizers to jax.example_libraries.optimizersNew features:
jax.lax.linalg.qdwh.jax.random.choice and jax.random.permutation now support multidimensional arrays and an optional axis argument
jax.random.choice and jax.random.permutation now support
multidimensional arrays and an optional axis argument (#8158)jax.numpy.take and jax.numpy.take_along_axis now require array-like inputs
(see #7737)New features:
jax.random.choice and jax.random.permutation now support
multidimensional arrays and an optional axis argument ({jax-issue}#8158)Breaking changes:
jax.numpy.take and jax.numpy.take_along_axis now require array-like inputs
(see {jax-issue}#7737)Nothing published for this version
The functions jax.ops.index_update, jax.ops.index_add etc. are deprecated and will be removed in a future JAX release. Please use the `.at` property o…
Static arguments to jax.pmap must now be hashable.
Unhashable static arguments have long been disallowed on jax.jit, but they
were still permitted on jax.pmap; jax.pmap compared unhashable static
arguments using object identity.
This behavior is a footgun, since comparing arguments using
object identity leads to recompilation each time the object identity
changes. Instead, we now ban unhashable arguments: if a user of jax.pmap
wants to compare static arguments by object identity, they can define
__hash__ and __eq__ methods on their objects that do that, or wrap their
objects in an object that has those operations with object identity
semantics. Another option is to use functools.partial to encapsulate the
unhashable static arguments into the function object.
jax.util.partial was an accidental export that has now been removed. Use
functools.partial from the Python standard library instead.
jax.ops.index_update, jax.ops.index_add etc. are
deprecated and will be removed in a future JAX release. Please use
the .at property on JAX arrays
instead, e.g., x.at[idx].set(y). For now, these functions produce a
DeprecationWarning.pmap is now the
default when using jaxlib 0.1.72 or newer. The feature can be disabled using
the --experimental_cpp_pmap flag (or JAX_CPP_PMAP environment variable).jax.numpy.unique now supports an optional fill_value argument ({jax-issue}#8121)Added jax.numpy.insert implementation (#7936 ).
New features:
jax.numpy.insert implementation (#7936 ).Breaking Changes
jax.api has been removed. Functions that were available as jax.api.*
were aliases for functions in jax.*; please use the functions in
jax.* instead.jax.partial, jax.lax.partial, and jax.util.partial were accidental
exports that have now been removed. Use functools.partial from the Python
standard library instead.TypeError; previously this silently
returned wrong results (#7925 ).jax.numpy functions now require array-like inputs, and will error
if passed a list (#7747 #7802 #7907 ).
See #7737 for a discussion of the rationale behind this change.jax.jit, jax.numpy.array always
stages the array it produces into the traced computation. Previously
jax.numpy.array would sometimes produce a on-device array, even under
a jax.jit decorator. This change may break code that used JAX arrays to
perform shape or index computations that must be known statically; the
workaround is to perform such computations using classic NumPy arrays
instead.jnp.ndarray is now a true base-class for JAX arrays. In particular, this
means that for a standard numpy array x, isinstance(x, jnp.ndarray) will
now return False (#7927).jax.api has been removed. Functions that were available as jax.api.*
were aliases for functions in jax.*; please use the functions in
jax.* instead.jax.partial, and jax.lax.partial were accidental exports that have now
been removed. Use functools.partial from the Python standard library
instead.TypeError; previously this silently
returned wrong results ({jax-issue}#7925).jax.numpy functions now require array-like inputs, and will error
if passed a list ({jax-issue}#7747 {jax-issue}#7802 {jax-issue}#7907).
See {jax-issue}#7737 for a discussion of the rationale behind this change.jax.jit, jax.numpy.array always
stages the array it produces into the traced computation. Previously
jax.numpy.array would sometimes produce a on-device array, even under
a jax.jit decorator. This change may break code that used JAX arrays to
perform shape or index computations that must be known statically; the
workaround is to perform such computations using classic NumPy arrays
instead.jnp.ndarray is now a true base-class for JAX arrays. In particular, this
means that for a standard numpy array x, isinstance(x, jnp.ndarray) will
now return False ({jax-issue}7927).jax.numpy.insert implementation ({jax-issue}#7936).jnp.poly* functions now require array-like inputs ({jax-issue}#7732)
jnp.poly* functions now require array-like inputs ({jax-issue}#7732)jnp.unique and other set-like operations now require array-like inputs
({jax-issue}#7662)GitHub commits .
Breaking Changes
jnp.poly* functions now require array-like inputs ( #7732 )
jnp.unique and other set-like operations now require array-like inputs ( #7662 )
Support for NumPy 1.17 has been dropped, per the deprecation policy. Please upgrade to a supported NumPy version.
Support for NumPy 1.17 has been dropped, per the deprecation policy. Please upgrade to a supported NumPy version.
The jit decorator has been added around the implementation of a number of
operators on JAX arrays. This speeds up dispatch times for common
operators such as +.
This change should largely be transparent to most users. However, there is
one known behavioral change, which is that large integer constants may now
produce an error when passed directly to a JAX operator
(e.g., x + 2**40). The workaround is to cast the constant to an
explicit type (e.g., np.float64(2**40)).
jnp.mean.
({jax-issue}#7317)#7613)Support for Python 3.6 has been dropped, per the deprecation policy. Please upgrade to a supported Python version.
Breaking changes:
backend argument to {py:func}jax.dlpack.from_dlpack has been
removed.New features:
jax.scipy.linalg.polar).Bug fixes:
axis value, or with an empty reduction dimension.
({jax-issue}#7196)Default to the older "stream_executor" CPU runtime for jaxlib <= 0.1.68 to work around #7229, which caused wrong outputs on CPU due to a concurrency p
jax.scipy.special.sph_harm.jax.grad,
{func}jax.value_and_grad, {func}jax.vjp, and
{func}jax.linear_transpose) support a parameter that indicates which named
axes should be summed over in the backward pass if they were broadcasted
over in the forward pass. This enables use of these APIs in a
non-per-example way inside maps (initially only
{func}jax.experimental.maps.xmap) ({jax-issue}#6950).GitHub commits .
Bug fixes:
Default to the older “stream_executor” CPU runtime for jaxlib <= 0.1.68 to work around #7229, which caused wrong outputs on CPU due to a concurrency problem.
New features:
New SciPy function jax.scipy.special.sph_harm() .
Reverse-mode autodiff functions ( jax.grad() , jax.value_and_grad() , jax.vjp() , and jax.linear_transpose() ) support a parameter that indicates which named axes should be summed over in the backward pass if they were broadcasted over in the forward pass. This enables use of these APIs in a non-per-example way inside maps (initially only jax.experimental.maps.xmap() ) ( #6950 ).
* GitHub commits.
Support for NumPy 1.16 has been dropped, per the deprecation policy.
New features:
jax2tf.convert supports inequalities and min/max for booleans
({jax-issue}#6956).jax.scipy.special.lpmn_values.Breaking changes:
Bug fixes:
jax2tf.call_tf(jax2tf.convert) ({jax-issue}#6947).GitHub commits .
New features:
#7042 Turned on TFRT CPU backend with significant dispatch performance improvements on CPU.
The jax2tf.convert() supports inequalities and min/max for booleans ( #6956 ).
New SciPy function jax.scipy.special.lpmn_values() .
Breaking changes:
Support for NumPy 1.16 has been dropped, per the deprecation policy .
Bug fixes:
Fixed bug that prevented round-tripping from JAX to TF and back: jax2tf.call_tf(jax2tf.convert) ( #6947 ).
The {func}jax2tf.convert now has support for pjit and sharded_jit.
New features:
jax2tf.convert now has support for pjit and sharded_jit.__tracebackhide__ is now enabled by
default in sufficiently recent versions of IPython.jax2tf.convert supports shape polymorphism even when the
unknown dimensions are used in arithmetic operations, e.g., jnp.reshape(-1)
({jax-issue}#6827).jax2tf.convert generates custom attributes with location information
in TF ops. The code that XLA generates after jax2tf
has the same location information as JAX/XLA.jax.scipy.special.lpmn.Bug fixes:
jax2tf.convert now ensures that it uses the same typing rules
for Python scalars and for choosing 32-bit vs. 64-bit computations
as JAX ({jax-issue}#6883).jax2tf.convert now scopes the enable_xla conversion parameter
properly to apply only during the just-in-time conversion
({jax-issue}#6720).jax2tf.convert now converts lax.dot_general using the
XlaDot TensorFlow op, for better fidelity w.r.t. JAX numerical precision
({jax-issue}#6717).jax2tf.convert now has support for inequality comparisons and
min/max for complex numbers ({jax-issue}#6892).GitHub commits .
New features:
The jax2tf.convert() now has support for pjit and sharded_jit .
A new configuration option JAX_TRACEBACK_FILTERING controls how JAX filters tracebacks.
A new traceback filtering mode using tracebackhide is now enabled by default in sufficiently recent versions of IPython.
The jax2tf.convert() supports shape polymorphism even when the unknown dimensions are used in arithmetic operations, e.g., jnp.reshape(-1) ( #6827 ).
The jax2tf.convert() generates custom attributes with location information in TF ops. The code that XLA generates after jax2tf has the same location information as JAX/XLA.
New SciPy function jax.scipy.special.lpmn() .
Bug fixes:
The jax2tf.convert() now ensures that it uses the same typing rules for Python scalars and for choosing 32-bit vs. 64-bit computations as JAX ( #6883 ).
The jax2tf.convert() now scopes the enable_xla conversion parameter properly to apply only during the just-in-time conversion ( #6720 ).
The jax2tf.convert() now converts lax.dot_general using the XlaDot TensorFlow op, for better fidelity w.r.t. JAX numerical precision ( #6717 ).
The jax2tf.convert() now has support for inequality comparisons and min/max for complex numbers ( #6892 ).
When combined with jaxlib 0.1.66, {func}jax.jit now supports static keyword arguments. A new static_argnames option has been added to specify keyword
New features:
jax.jit now supports static
keyword arguments. A new static_argnames option has been added to specify
keyword arguments as static.jax.nonzero has a new optional size argument that allows it to
be used within jit ({jax-issue}#6501)jax.numpy.unique now supports the axis argument ({jax-issue}#6532).jax.experimental.host_callback.call now supports pjit.pjit ({jax-issue}#6569).jax.scipy.linalg.eigh_tridiagonal that computes the
eigenvalues of a tridiagonal matrix. Only eigenvalues are supported at
present.UnfilteredStackTrace exception
containing the original trace as the __cause__ of the filtered exception.
Filtered stack traces now also work with Python 3.6.__cause__ of
the exception a JaxStackTraceBeforeTransformation object that contains the
stack trace that created the original operation in the forward pass.
Requires jaxlib 0.1.66.Breaking changes:
host_id --> {func}~jax.process_indexhost_count --> {func}~jax.process_counthost_ids --> range(jax.process_count())~jax.local_devices has been renamed from
host_id to process_index.jax.jit other than the function are now marked as
keyword-only. This change is to prevent accidental breakage when arguments
are added to jit.Bug fixes:
jax2tf.convert now works in presence of gradients for functions
with integer inputs ({jax-issue}#6360).jax2tf.call_tf when used with captured
tf.Variable ({jax-issue}#6572).New profiling APIs: {func}jax.profiler.start_trace, {func}jax.profiler.stop_trace, and {func}jax.profiler.trace
jax.profiler.start_trace,
{func}jax.profiler.stop_trace, and {func}jax.profiler.tracejax.lax.reduce is now differentiable.TraceContext --> {func}~jax.profiler.TraceAnnotationStepTraceContext --> {func}~jax.profiler.StepTraceAnnotationtrace_function --> {func}~jax.profiler.annotate_functionint64 value will now lead to an overflow
in all cases, rather than being silently converted to uint64 in some cases ({jax-issue}#6047).int32 will now lead to an
OverflowError rather than having their value silently truncated.host_callback now supports empty arrays in arguments and results ({jax-issue}#6262).jax.random.randint clips rather than wraps of out-of-bounds limits, and can now generate
integers in the full range of the specified dtype ({jax-issue}#5868)GitHub commits .
New features
New profiling APIs: jax.profiler.start_trace() , jax.profiler.stop_trace() , and jax.profiler.trace()
jax.lax.reduce() is now differentiable.
Breaking changes:
The minimum jaxlib version is now 0.1.64.
Some profiler APIs names have been changed. There are still aliases, so this should not break existing code, but the aliases will eventually be removed so please change your code.
TraceContext –> TraceAnnotation()
StepTraceContext –> StepTraceAnnotation()
trace_function –> annotate_function()
Omnistaging can no longer be disabled. See omnistaging for more information.
Python integers larger than the maximum int64 value will now lead to an overflow in all cases, rather than being silently converted to uint64 in some cases ( #6047 ).
Outside X64 mode, Python integers outside the range representable by int32 will now lead to an OverflowError rather than having their value silently truncated.
Bug fixes:
host_callback now supports empty arrays in arguments and results ( #6262 ).
jax.random.randint() clips rather than wraps of out-of-bounds limits, and can now generate integers in the full range of the specified dtype ( #5868 )
The minimum jaxlib version is now 0.1.62.
New features:
Bug fixes:
jax.flatten_util.ravel_pytree to handle integer dtypes.enum.IntEnumsBreaking changes:
{func}jax.scipy.stats.chi2 is now available as a distribution with logpdf and pdf methods.
jax.scipy.stats.chi2 is now available as a distribution with logpdf and pdf methods.jax.scipy.stats.betabinom is now available as a distribution with logpmf and pmf methods.jax.experimental.jax2tf.call_tf to call TensorFlow functions
from JAX ({jax-issue}#5627)
and README).lax.pad to support batching of the padding values.jax.numpy.take properly handles negative indices ({jax-issue}#5768)jnp.bfloat16(1) + 0.1 * jnp.arange(10)
previously returned a float64 array, and now returns a bfloat16 array.
JAX's type promotion behavior is described at {ref}type-promotion.jax.numpy.linspace now computes the floor of integer values, i.e.,
rounding towards -inf rather than 0. This change was made to match NumPy
1.20.0.jax.numpy.i0 no longer accepts complex numbers. Previously the
function computed the absolute value of complex arguments. This change was
made to match the semantics of NumPy 1.20.0.jax.numpy functions no longer accept tuples or lists in place
of array arguments: {func}jax.numpy.pad, :funcjax.numpy.ravel,
{func}jax.numpy.repeat, {func}jax.numpy.reshape.
In general, {mod}jax.numpy functions should be used with scalars or array arguments.GitHub commits .
New features:
jax.scipy.stats.chi2() is now available as a distribution with logpdf and pdf methods.
jax.scipy.stats.betabinom() is now available as a distribution with logpmf and pmf methods.
Added jax.experimental.jax2tf.call_tf() to call TensorFlow functions from JAX ( #5627 ) and README ).
Extended the batching rule for lax.pad to support batching of the padding values.
Bug fixes:
jax.numpy.take() properly handles negative indices ( #5768 )
Breaking changes:
JAX’s promotion rules were adjusted to make promotion more consistent and invariant to JIT. In particular, binary operations can now result in weakly-typed values when appropriate. The main user-visible effect of the change is that some operations result in outputs of different precision than before; for example the expression jnp.bfloat16(1) + 0.1 * jnp.arange(10) previously returned a float64 array, and now returns a bfloat16 array. JAX’s type promotion behavior is described at Type promotion semantics .
jax.numpy.linspace() now computes the floor of integer values, i.e., rounding towards -inf rather than 0. This change was made to match NumPy 1.20.0.
jax.numpy.i0() no longer accepts complex numbers. Previously the function computed the absolute value of complex arguments. This change was made to match the semantics of NumPy 1.20.0.
Several jax.numpy functions no longer accept tuples or lists in place of array arguments: jax.numpy.pad() , :func jax.numpy.ravel , jax.numpy.repeat() , jax.numpy.reshape() . In general, jax.numpy functions should be used with scalars or array arguments.
Extend the {mod}jax.experimental.loops module with support for pytrees. Improved error checking and error messages.
jax.experimental.loops module with support for pytrees. Improved
error checking and error messages.jax.experimental.enable_x64 and {func}jax.experimental.disable_x64.
These are context managers which allow X64 mode to be temporarily enabled/disabled
within a session.jax.ops.segment_sum now drops segment IDs that are out of range rather
than wrapping them into the segment ID space. This was done for performance
reasons.GitHub commits .
New features:
Extend the jax.experimental.loops module with support for pytrees. Improved error checking and error messages.
Add jax.experimental.enable_x64() and jax.experimental.disable_x64() . These are context managers which allow X64 mode to be temporarily enabled/disabled within a session.
Breaking changes:
jax.ops.segment_sum() now drops segment IDs that are out of range rather than wrapping them into the segment ID space. This was done for performance reasons.
Removed support for kwargs for {func}jax.experimental.host_callback.id_tap. (This support has been deprecated for a few months.)
jax.closure_convert for use with higher-order custom
derivative functions. ({jax-issue}#5244)jax.experimental.host_callback.call to call a custom Python
function on the host and return a result to the device computation.
({jax-issue}#5243)jax.numpy.arccosh now returns the same branch as numpy.arccosh for
complex inputs ({jax-issue}#5156)host_callback.id_tap now works for jax.pmap also. There is an
optional parameter for id_tap and id_print to request that the
device from which the value is tapped be passed as a keyword argument
to the tap function ({jax-issue}#5182).jax.numpy.pad now takes keyword arguments. Positional argument constant_values
has been removed. In addition, passing unsupported keyword arguments raises an error.jax.experimental.host_callback.id_tap ({jax-issue}#5243):
kwargs for {func}jax.experimental.host_callback.id_tap.
(This support has been deprecated for a few months.)jax.experimental.host_callback.id_print
to use '(' instead of '['.jax.experimental.host_callback.id_print in presence of JVP
to print a pair of primal and tangent. Previously, there were two separate
print operations for the primals and the tangent.host_callback.outfeed_receiver has been removed (it is not necessary,
and was deprecated a few months ago).inf, analogous to that for NaN ({jax-issue}#5224).indexing of JAX arrays with non-tuple sequences now raises a TypeError. This type of indexing has been deprecated in Numpy since v1.16, and in JAX sin…
jax.device_put_replicatedjax.experimental.sharded_jitjax.numpy.linalg.eigjax.pmapjax.numpy.linalg.slogdetjax.numpy.sinc at zerojax.experimental.optix has been deleted, in favor of the standalone
optax Python package.TypeError. This type of indexing
has been deprecated in Numpy since v1.16, and in JAX since v0.2.4.
See {jax-issue}#4564.Add support for shape-polymorphic tracing for the jax.experimental.jax2tf converter. See README.md.
New Features:
Breaking change cleanup
Raise an error on non-hashable static arguments for jax.jit and xla_computation. See cb48f42.
Improve consistency of type promotion behavior ({jax-issue}#4744):
jnp.float32(1) + 1j now returns complex64, where previously
it returned complex128.jnp.result_type(jnp.uint64, jnp.int64, jnp.float16) and
jnp.result_type(jnp.float16, jnp.uint64, jnp.int64) both return float16, where previously
the first returned float64 and the second returned float16.The contents of the (undocumented) jax.lax_linalg linear algebra module
are now exposed publicly as jax.lax.linalg.
jax.random.PRNGKey now produces the same results in and out of JIT compilation
({jax-issue}#4877).
This required changing the result for a given seed in a few particular cases:
jax_enable_x64=False, negative seeds passed as Python integers now return a different result
outside JIT mode. For example, jax.random.PRNGKey(-1) previously returned
[4294967295, 4294967295], and now returns [0, 4294967295]. This matches the behavior in JIT.int64 outside JIT now result in an OverflowError
rather than a TypeError. This matches the behavior in JIT.To recover the keys returned previously for negative integers with jax_enable_x64=False
outside JIT, you can use:
key = random.PRNGKey(-1).at[0].set(0xFFFFFFFF)
DeviceArray now raises RuntimeError instead of ValueError when trying
to access its value while it has been deleted.
Ensure that check_jaxpr does not perform FLOPS. See {jax-issue}#4650.
check_jaxpr does not perform FLOPS. See {jax-issue}#4650.Indexing with non-tuple sequences is now deprecated, following a similar deprecation in Numpy. In a future release, this will result in a TypeError. S…
Improvements:
remat to jax.experimental.host_callback. See {jax-issue}#4608.Deprecations
#4564.The reason for another release so soon is we need to temporarily roll back a new jit fastpath while we look into a performance degradation
* GitHub commits.
As a benefit of omnistaging, the host_callback functions are executed (in program order) even if the result of the {py:func}jax.experimental.host_call
jax.experimental.host_callback.id_print/
{py:func}jax.experimental.host_callback.id_tap is not used in the computation.GitHub commits .
Improvements:
As a benefit of omnistaging, the host_callback functions are executed (in program order) even if the result of the jax.experimental.host_callback.id_print() / jax.experimental.host_callback.id_tap() is not used in the computation.
Omnistaging on by default. See {jax-issue}#3370 and omnistaging
#3370 and
omnistagingGitHub commits .
Improvements:
Omnistaging on by default. See #3370 and omnistaging
New simplified interface for {py:func}jax.experimental.host_callback.id_tap
jax.experimental.host_callback.id_tap (#4101)Breaking changes:
New simplified interface for jax.experimental.host_callback.id_tap() (#4101)
Support for NumPy 1.18 has been dropped, per the deprecation policy. Please upgrade to a supported NumPy version.
Your coding agent can read these notes before it upgrades. Set up the MCP server →