NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #1265 most downloaded on PyPI
Differentiate, compile, and transform Numpy code.
Last release 17 days ago
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
* New features: * Support for python 3.10.
Enabled peer-to-peer copies between GPUs. Previously, GPU copies were bounced via the host, which is usually slower.
One column per quarter.
Multiple cuDNN versions are now supported for jaxlib GPU cuda11 wheels.
Multiple cuDNN versions are now supported for jaxlib GPU cuda11 wheels.
Breaking changes:
The install commands for GPU jaxlib are as follows:
pip install --upgrade pip
# Installs the wheel compatible with CUDA 11 and cuDNN 8.2 or newer.
pip install --upgrade "jax[cuda]" -f https://storage.googleapis.com/jax-releases/jax_releases.html
# Installs the wheel compatible with Cuda 11 and cudnn 8.2 or newer.
pip install jax[cuda11_cudnn82] -f https://storage.googleapis.com/jax-releases/jax_releases.html
# Installs the wheel compatible with Cuda 11 and cudnn 8.0.5 or newer.
pip install jax[cuda11_cudnn805] -f https://storage.googleapis.com/jax-releases/jax_releases.html
Support for CUDA 10.2 and CUDA 10.1 has been dropped. Jaxlib now supports CUDA 11.1+.
Support for CUDA 11.0 and CUDA 10.1 has been dropped. Jaxlib now supports CUDA 10.2 and CUDA 11.1+.
Support for Python 3.6 has been dropped, per the deprecation policy. Please upgrade to a supported Python version.
Support for Python 3.6 has been dropped, per the deprecation policy. Please upgrade to a supported Python version.
Support for NumPy 1.17 has been dropped, per the deprecation policy. Please upgrade to a supported NumPy version.
The host_callback mechanism now uses one thread per local device for making the calls to the Python callbacks. Previously there was a single thread for all devices. This means that the callbacks may now be called interleaved. The callbacks corresponding to one device will still be called in sequence.
Fix bugs in TFRT CPU backend that results in incorrect results.
Fixed bug in TFRT CPU backend that gets nans when transfer TPU buffer to CPU.
Support for reduction over subsets of a pmapped axis using axis_index_groups {jax-issue}#2382.
axis_index_groups
{jax-issue}#2382.#3006).jax.numpy has been
tightened. This may break code that was making use of names that were
previously exported accidentally.CUDA 11.1 wheels are now supported on all CUDA 11 versions 11.1 or higher.
CUDA 11.1 wheels are now supported on all CUDA 11 versions 11.1 or higher.
NVidia now promises compatibility between CUDA minor releases starting with CUDA 11.1. This means that JAX can release a single CUDA 11.1 wheel that is compatible with CUDA 11.2 and 11.3.
There is no longer a separate jaxlib release for CUDA 11.2 (or higher); use the CUDA 11.1 wheel for those versions (cuda111).
Jaxlib now bundles libdevice.10.bc in CUDA wheels. There should be no need
to point JAX to a CUDA installation to find this file.
Added automatic support for static keyword arguments to the {func}jit
implementation.
Added support for pretransformation exception traces.
Initial support for pruning unused arguments from {func}jit -transformed
computations.
Pruning is still a work in progress.
Improved the string representation of {class}PyTreeDef objects.
Added support for XLA's variadic ReduceWindow.
jit transformed functions.Differentiation of determinants of singular matrices {jax-issue}#2809.
#2809.odeint differentiation with respect to time of ODEs with
time-dependent dynamics {jax-issue}#2817,
also add ODE CI testing.lax_linalg.qr differentiation
{jax-issue}#2867.GitHub commits .
New features:
Differentiation of determinants of singular matrices #2809 .
Bug fixes:
Fix odeint() differentiation with respect to time of ODEs with time-dependent dynamics #2817 , also add ODE CI testing.
Fix lax_linalg.qr() differentiation #2867 .
Add syntactic sugar for functional indexed updates {jax-issue}#2684.
#2684.jax.numpy.linalg.multi_dot {jax-issue}#2726.jax.numpy.unique {jax-issue}#2760.jax.numpy.rint {jax-issue}#2724.jax.numpy.rint {jax-issue}#2724.jax.experimental.jet.logaddexp and {func}logaddexp2 differentiation at zero {jax-issue}#2107.jit
{jax-issue}#2719.lax.while_loop
{jax-issue}#2129.GitHub commits .
New features:
Add syntactic sugar for functional indexed updates #2684 .
Add jax.numpy.linalg.multi_dot() #2726 .
Add jax.numpy.unique() #2760 .
Add jax.numpy.rint() #2724 .
Add jax.numpy.rint() #2724 .
Add more primitive rules for jax.experimental.jet() .
Bug fixes:
Fix logaddexp() and logaddexp2() differentiation at zero #2107 .
Improve memory usage in reverse-mode autodiff without jit() #2719 .
Better errors:
Improves error message for reverse-mode differentiation of lax.while_loop() #2129 .
Added jax.custom_jvp and jax.custom_vjp from {jax-issue}#2026, see the tutorial notebook. Deprecated jax.custom_transforms and removed it from the doc…
jax.custom_jvp and jax.custom_vjp from {jax-issue}#2026, see the tutorial notebook. Deprecated jax.custom_transforms and removed it from the docs (though it still works).scipy.sparse.linalg.cg {jax-issue}#2566.#2591.jax.numpy.isclose handle nan and inf correctly {jax-issue}#2501.jax.experimental.jet {jax-issue}#2537.jax.experimental.stax.BatchNorm when scale/center isn't provided.jax.numpy.einsum {jax-issue}#2512.jax.numpy.cumsum and jax.numpy.cumprod in terms of a parallel prefix scan {jax-issue}#2596 and make reduce_prod differentiable to arbitrary order {jax-issue}#2597.batch_group_count to conv_general_dilated {jax-issue}#2635.test_util.check_grads {jax-issue}#2656.callback_transform {jax-issue}#2665.rollaxis, convolve/correlate 1d & 2d, copysign,
trunc, roots, and quantile/percentile interpolation options.GitHub commits .
Added jax.custom_jvp and jax.custom_vjp from #2026 , see the tutorial notebook . Deprecated jax.custom_transforms and removed it from the docs (though it still works).
Add scipy.sparse.linalg.cg #2566 .
Changed how Tracers are printed to show more useful information for debugging #2591 .
Made jax.numpy.isclose handle nan and inf correctly #2501 .
Added several new rules for jax.experimental.jet #2537 .
Fixed jax.experimental.stax.BatchNorm when scale / center isn’t provided.
Fix some missing cases of broadcasting in jax.numpy.einsum #2512 .
Implement jax.numpy.cumsum and jax.numpy.cumprod in terms of a parallel prefix scan #2596 and make reduce_prod differentiable to arbitrary order #2597 .
Add batch_group_count to conv_general_dilated #2635 .
Add docstring for test_util.check_grads #2656 .
Add callback_transform #2665 .
Implement rollaxis , convolve / correlate 1d & 2d, copysign , trunc , roots , and quantile / percentile interpolation options.
jaxlib wheels are now built to require AVX instructions on x86-64 machines by default. If you want to use JAX on a machine that doesn't support AVX, y
--target_cpu_features flag
to build.py. --target_cpu_features also replaces
--enable_march_native.Fixes Python 3.5 support. This will be the last JAX or jaxlib release that supports Python 3.5.
Fixed a memory leak when converting CPU DeviceArrays to NumPy arrays. The memory leak was present in jaxlib releases 0.1.58 and 0.1.59.
bool, int8, and uint8 are now considered safe to cast to
bfloat16 NumPy extension type.The minimum jaxlib version is now 0.1.38.
Breaking changes
Jaxpr by removing the Jaxpr.freevars and
Jaxpr.bound_subjaxprs. The call primitives (xla_call, xla_pmap,
sharded_call, and remat_call) get a new parameter call_jaxpr with a
fully-closed (no constvars) jaxpr. Also, added a new field call_primitive
to primitives.New features:
grad) of lax.cond, making it
now differentiable in both modes ({jax-issue}#2091)__cuda_array_interface__, which is another
zero-copy protocol for sharing GPU arrays with other libraries such as CuPy
and Numba.Fixed a bug that meant JAX sometimes return platform-specific types (e.g., np.cint) instead of standard types (e.g., np.int32).
np.cint) instead of standard types (e.g., np.int32). (#4903)is_leaf predicate to {func}pytree.flatten.Fixed manylinux2010 compliance issues in GPU wheels.
Nothing published for this version
Fix bug in DLPackManagedTensorToBuffer
Nothing published for this version
Nothing published for this version
* Update XLA.
Add new runtime support for host_callback.
Drop support for CUDA 9.2 (we only maintain support for the last four CUDA versions.)
Fix build issue that could result in slow compiles ( )
Bug fixes:
Fix build issue that could result in slow compiles ( tensorflow/tensorflow )
Adds support for fast traceback collection.
np.nextafter for bfloat16 types.tanh accuracy on GPU.* Fixes crash for outfeed.
Fixes crash for linear algebra functions on Mac OS X (#432).
Fixes segfault: {jax-issue}#2755
#2755Fixes segfault: #2755
Plumb is_stable option on Sort HLO through to Python.
Fixes a bug where if multiple GPUs of different models were present, JAX would only compile programs suitable for the first GPU.
batch_group_count convolutions.Fixed a performance regression for Resnet-50 on GPU.
jaxlib 0.1.41 broke cloud TPU support due to an API incompatibility. This release fixes it again.
Nothing published for this version
Adds experimental support in Jaxlib for TensorFlow profiler, which allows tracing of CPU and GPU computations from TensorBoard.
* Updates XLA.
CUDA 9.0 is no longer supported.
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
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 →