NewYour coding agent can read the release notes before it upgrades.Set up the MCP server →
PyPI · #1286 most downloaded on PyPI
XLA library for JAX
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 42 of 42 stable releases
1 version withdrawn
withdrawn after publishing
3 years old
43 releases · first in 2023
One column per quarter.
Removed deprecated jax.experimental.shard_alike . Use explicit sharding mode instead (see sharding ).
jax.numpy.minmax (and jnp.minmax), which returns both thejax.lax.log2 and primitive jax.lax.log2_p, makinglog2 a first-class primitive in JAX (jax.numpy.log2 now lowers viajax.lax.log2).jax.lax.one_minus_square primitive to accurately compute1 - x^2 near $\pm 1$ and provide accurate derivatives near $0$.jax.export.symbolic_dim_bounds for querying conservativefrozendict support to JAX pytrees for Python 3.15 (PEP 814).jax.distributed.initialize can now secure the coordination servicemtls_cert_file, mtls_key_file,mtls_ca_file, mtls_peer_uri_prefix and verify_secure_credentialsJAX_MTLS_CERT_FILE, JAX_MTLS_KEY_FILE,JAX_MTLS_CA_FILE, JAX_MTLS_PEER_URI_PREFIX andJAX_DISTRIBUTED_VERIFY_SECURE_CREDENTIALS environment variables).jax.distributed.initialize (#40512).TPU_PROCESS_ADDRESSES_PATH in GKE TPU clusters.jax.random.generalized_normal's p parameter type fromfloat to RealArray, allowing array-valued shape parameters (#40126).exclude_argnames argument to jax.experimental.program_order.geqrf, orgqr/ungqr, ormqr/unmqr), LU decomposition (getrf),syevd/heevd), SVDgesvd), and hybrid solver kernels (geqp3, eig)jaxlib wheels now ship C++ FFI extension headers (collectives.h,record.h) to support out-of-tree plugins (#40333).jax.experimental.shard_alike. Use explicit shardingsharding).jax.sharding.Mesh construction by avoidingjaxlib for free-threaded Pythoninline=True in jax.jit now corresponds tojax.Inline.JAX_LATE instead of jax.Inline.JAX_EARLY.jax.numpy.fft.irfftn,jax.numpy.fft.irfft2 and jax.lax.fft with FftType.IRFFT)jax.numpy.cumsum)jax.numpy.tri now returns an array with the default float dtypedtype argument is not specified. Previously it always returnedfloat32 (#40242).jax.numpy.unique with axis specified now matches NumPy's outputjax.numpy.log2 by pre-computing the1 / log(2) constant factor (#40430).out_sharding parameter to jax.numpy.histogram.jax.remat's prevent_cse argument signature to acceptbool | Sequence[bool], matching jax.checkpoint.jax.experimental.checkify error code assignment deterministic.jax.numpy.arccosh and jax.lax.acoshjax.lax.bessel_i0e and jax.lax.bessel_i1e at 0.0jax.numpy.linalg.eigh gradients producing NaN or incorrectjax.numpy.sinc now uses a Taylor series near zero, givingsin(πx)/(πx) quotient suffered catastrophic cancellation near zerojax.numpy.linalg.cond returned NaN instead ofp is None or 2, matching NumPyjax.numpy.histogram crashing on empty arraysjax.numpy.intersect1d and jax.numpy.setxor1d withsize=0, which previously raised a ValueError; they now return emptyjax.numpy.setdiff1d raising an IndexError when called withsize=0 on non-empty inputs; it now returns an empty array.jax.scipy.linalg.cholesky andjax.numpy.linalg.cholesky with symmetrize_input=False wherejax.numpy.median on an input that is empty along thegather;ValueError.jax.nn.initializers.variance_scaling for zero-size inputsjax.custom_root tangents when auxiliary values arejax.debug.visualize_array_shardingjax.lax.min and jax.lax.max to notjax.export whenref_addupdate) on indexedReshapeTransform views.New features
jax.numpy.minmax (and jnp.minmax), which returns both the
minimum and maximum of an array, matching NumPy 2.3+ ({jax-issue}#40089).jax.lax.log2 and primitive {data}jax.lax.log2_p, making
log2 a first-class primitive in JAX ({func}jax.numpy.log2 now lowers via
jax.lax.log2).jax.lax.one_minus_square primitive to accurately compute
1 - x^2 near $\pm 1$ and provide accurate derivatives near $0$.jax.export.symbolic_dim_bounds for querying conservative
bounds on symbolic dimension expressions ({jax-issue}#40006).frozendict support to JAX pytrees for Python 3.15 (PEP 814).jax.distributed.initialize can now secure the coordination service
with mutual TLS via the new mtls_cert_file, mtls_key_file,
mtls_ca_file, mtls_peer_uri_prefix and verify_secure_credentials
arguments (or the JAX_MTLS_CERT_FILE, JAX_MTLS_KEY_FILE,
JAX_MTLS_CA_FILE, JAX_MTLS_PEER_URI_PREFIX and
JAX_DISTRIBUTED_VERIFY_SECURE_CREDENTIALS environment variables).jax.distributed.initialize ({jax-issue}#40512).TPU_PROCESS_ADDRESSES_PATH in GKE TPU clusters.jax.random.generalized_normal's p parameter type from
float to RealArray, allowing array-valued shape parameters ({jax-issue}#40126).exclude_argnames argument to {func}jax.experimental.program_order.geqrf, orgqr/ungqr, ormqr/unmqr), LU decomposition (getrf),
symmetric/Hermitian eigenvalue decomposition (syevd/heevd), SVD
(gesvd), and hybrid solver kernels (geqp3, eig)
({jax-issue}#40000, {jax-issue}#40186, {jax-issue}#40543).jaxlib wheels now ship C++ FFI extension headers (collectives.h,
record.h) to support out-of-tree plugins ({jax-issue}#40333).Breaking changes
jax.experimental.shard_alike. Use explicit sharding
mode instead (see {ref}jax-201-sharding).Changes
jax.sharding.Mesh construction by avoiding
redundant device array allocations and copies.jaxlib for free-threaded Python
(Python 3.13t, 3.14t, 3.15t).inline=True in {func}jax.jit now corresponds to
{attr}jax.Inline.JAX_LATE instead of {attr}jax.Inline.JAX_EARLY.jax.numpy.fft.irfftn,
{func}jax.numpy.fft.irfft2 and {func}jax.lax.fft with FftType.IRFFT)
are again lowered to a single C2R transform, as before JAX 0.10.0, instead
of an IFFT over the outer axes and a 1-D IRFFT with two transposes. The
input is first made Hermitian-symmetric along the outer axes, which does
not change the result under NumPy's convention (only the last axis is
assumed symmetric), so results are unchanged while the transform is
~1.4x faster at typical sizes.jax.numpy.cumsum)
on GPU, improving performance.jax.numpy.tri now returns an array with the default float dtype
when the dtype argument is not specified. Previously it always returned
float32 ({jax-issue}#40242).jax.numpy.unique with axis specified now matches NumPy's output
shape for arrays that are empty along the given axis, instead of
fabricating a phantom slice for fully-empty inputs.jax.numpy.log2 by pre-computing the
1 / log(2) constant factor ({jax-issue}#40430).out_sharding parameter to {func}jax.numpy.histogram.jax.remat's prevent_cse argument signature to accept
bool | Sequence[bool], matching {func}jax.checkpoint.jax.experimental.checkify error code assignment deterministic.Bug fixes
jax.numpy.arccosh and {func}jax.lax.acosh
gradients for large inputs ({jax-issue}#40643, {jax-issue}#40634).jax.lax.bessel_i0e and {func}jax.lax.bessel_i1e at 0.0
({jax-issue}#40640, {jax-issue}#40635).jax.numpy.linalg.eigh gradients producing NaN or incorrect
values for large eigenvalues ({jax-issue}#40149, {jax-issue}#40141).jax.numpy.sinc now uses a Taylor series near zero, giving
accurate derivatives of all orders. Previously, autodiff of the
sin(πx)/(πx) quotient suffered catastrophic cancellation near zero
({jax-issue}#34139, {jax-issue}#10750).jax.numpy.partition, {func}jax.numpy.argpartition and
{func}jax.numpy.top_k with mode='smallest' on signed integer arrays
containing the smallest representable value, which was treated as the
largest element and dropped from the result.jax.numpy.partition and {func}jax.numpy.argpartition now accept
boolean arrays, as {func}jax.numpy.top_k already did.jax.numpy.linalg.cond returned NaN instead of
infinity for singular matrices when p is None or 2, matching NumPy
and the other norms.jax.numpy.histogram crashing on empty arrays
({jax-issue}#40025, {jax-issue}#40020).jax.numpy.intersect1d and {func}jax.numpy.setxor1d with
size=0, which previously raised a ValueError; they now return empty
arrays of the natural result dtype.jax.numpy.setdiff1d raising an IndexError when called with
size=0 on non-empty inputs; it now returns an empty array.jax.scipy.linalg.cholesky and
{func}jax.numpy.linalg.cholesky with symmetrize_input=False where
non-zero gradients leaked into the unused triangle of the input matrix
({jax-issue}#40421).jax.numpy.median on an input that is empty along the
reduction axis, which previously raised an internal error from gather;
it now raises a ValueError.jax.nn.initializers.variance_scaling for zero-size inputs
({jax-issue}#35096).jax.custom_root tangents when auxiliary values are
integer-typed ({jax-issue}#39913, {jax-issue}#24295).jax.debug.visualize_array_sharding
({jax-issue}#39922, {jax-issue}#25695).jax.lax.min and {func}jax.lax.max to not
depend on bitwise equivalence between forward and backward pass results
({jax-issue}#40578).jax.export when
even-powered factor bounds cross zero or zero factors are paired with
infinite bounds ({jax-issue}#40054).ref_addupdate) on indexed
ReshapeTransform views.#40389).Deprecations
pl.reciprocal was moved into {mod}jax.experimental.pallas.tpu.
Accessing it via {mod}jax.experimental.pallas is deprecated.Removals
pl.dot. Use {func}jax.numpy.dot,
{func}jax.numpy.einsum or the @ operator instead in TPU or MGPU kernels.pl.debug_checks_enabled. Use
pl.enable_debug_checks.value instead.pltpu.HOST. Use pl.HOST instead.scrach_shapes and out_shapes
arguments from {func}jax.experimental.pallas.mosaic_gpu.kernel. Use
scratch_types and out_type instead.### Mosaic GPU
Changes
jax.experimental.pallas.mosaic_gpu.try_cluster_cancel now takes a
collective_axes argument. It is now not, for example, collective in the
thread axis unless that is included in collective_axes. Users of
{func}jax.experimental.pallas.mosaic_gpu.dynamic_scheduling_loop that fail
to pass complete thread_axis and/or cluster_axes arguments will now see
multiple cluster cancellations (these arguments have always been mandatory,
but it may previously have worked without, albeit non-deterministically!).New features
MOSAIC_GPU_DUMP_RESOURCES environment variable, which dumps
the estimated resources (e.g. TMEM columns, SMEM bytes) when compiling a
model. Like the other dump variables, it is also enabled by
MOSAIC_GPU_DUMP_TO.MOSAIC_GPU_DUMP_CONSTRAINT_SYSTEM environment variable, which
dumps the constraint system built during layout inference. This is useful
for debugging layout inference failures. Like the other dump variables, it
is also enabled by MOSAIC_GPU_DUMP_TO.jax.experimental.pallas.mosaic_gpu.store, the counterpart to
{func}jax.experimental.pallas.mosaic_gpu.load. Like load, it takes an
optimized argument that controls whether a conflict-free transfer is
required.Deprecations
jax.experimental.pallas.mosaic_gpu.transpose_ref is
deprecated. Use ref.transpose(...) directly instead.jax.experimental.pallas.program_id and
{func}jax.experimental.pallas.num_programs in Pallas MGPU kernels is
deprecated. Use {func}jax.lax.axis_index and {func}jax.lax.axis_size
instead.Removals
jax.experimental.pallas.mosaic_gpu.transform_ref. It is
only marginally useful in the presence of transform inference. If you find
that transform inference is insufficient, please file a bug.jax.experimental.pallas.pallas_call. Mosaic GPU kernels can now
only be defined via {func}jax.experimental.pallas.mosaic_gpu.kernel.New features
jax.experimental.pallas.tpu.annotate to attach memory access
assumptions (no_store, no_bank_conflict, no_hazard, and
no_hazard_no_deps) to VMEM/SMEM references, allowing kernels to
override compiler scheduling, bundle packing, and memory dependency analysis.jax_pallas_auto_assign_collective_ids_limit config flag to allow
configuring the limit for auto-assigned collective IDs.jax.experimental.pallas.tpu.get_barrier_semaphore now accepts an
optional hashable tag argument, assigning barrier semaphores automatically
to kernels sharing the same tag.The fields in_shardings_hlo and out_shardings_hlo of jax.export.Exported have been deprecated for a while. Now accessing them raises a warning. Use in…
New features
--jax_export_deserialize_expired_versions tojax.numpy.top_k, which implements numpy.top_k, added inBreaking changes
exec_time_optimization_effort and memory_fitting_effort flags have beenEffortLevel enum.Deprecations
in_shardings_hlo and out_shardings_hlo ofjax.export.Exported have been deprecated for a while. Now accessing themin_shardings_jax and out_shardings_jax instead.Changes
jax.nn.dot_product_attention with implementation='cudnn') nomask, whose gradient no caller can request. Bias gradientsbias or a non-boolean mask are unchangedjax.numpy.meshgrid, jax.numpy.ogrid, andjax.numpy.broadcast_arrays now return tuples rather than listsjax.grad or jax.value_and_grad rejects a function withoutput.sum()), using jax.jacobian, orjax.lax.dynamic_slice,jax.lax.dynamic_update_slice, or jax.ds, and shows tracerBug fixes
jax.numpy.linalg.det and jax.numpy.linalg.slogdet now use ajax.nn.dot_product_attention with implementation='cudnn') nowmask. Previously jax.jacobian, jax.vmap with partialin_axes, and jax.vmap of a VJP or of jax.grad failed with aTypeError (#38495).jax.vmap of fp8 cuDNN fused attention now works: its batching rulesNotImplementedError.jax_compiler_enable_remat_pass to False now addsrematerialization to the set of disabled XLA passes instead ofXLA_FLAGS=--xla_disable_hlo_passes=... stay disabledjax.numpy.split, jax.numpy.array_split, and thehsplit/vsplit/dsplit variants once again accept negative entries inindices_or_sections, resolving them against the axis size as NumPy doesValueError: Sizes passed to split must be nonnegative.jax.lax.scan to only check .matShapedArrayjax.lax.reshapejax.tree_util.flatten_one_level_with_keys for namedtuple_get_prime_factors in jax.experimental.mesh_utilsNew features
--jax_export_deserialize_expired_versions to
temporarily bypass the error check.
See https://docs.jax.dev/en/latest/export/export.html#compatibility-guarantees.jax.numpy.top_k, which implements {func}numpy.top_k, added in
in NumPy v2.6.0 ({jax-issue}#39729).Breaking changes
exec_time_optimization_effort and memory_fitting_effort flags have been
removed in favor of the EffortLevel enum.Deprecations
in_shardings_hlo and out_shardings_hlo of
jax.export.Exported have been deprecated for a while. Now accessing them
raises a warning. Use in_shardings_jax and out_shardings_jax instead.Changes
jax.nn.dot_product_attention with implementation='cudnn') no
longer computes a bias gradient when the only attention bias comes from
a boolean mask, whose gradient no caller can request. Bias gradients
for an explicit bias or a non-boolean mask are unchanged
({jax-issue}#34685).jax.numpy.meshgrid, {obj}jax.numpy.ogrid, and
{func}jax.numpy.broadcast_arrays now return tuples rather than lists
in order to align with NumPy>2.0 and the Array API specification.
({jax-issue}#39783, {jax-issue}#39789, {jax-issue}#39802)jax.grad or {func}jax.value_and_grad rejects a function with
a non-scalar output, the error message now suggests reducing the output
to a scalar (e.g. with output.sum()), using {func}jax.jacobian, or
reshaping size-1 outputs ({jax-issue}#2303).jax.lax.dynamic_slice,
{func}jax.lax.dynamic_update_slice, or jax.ds, and shows tracer
provenance ({jax-issue}#7222).#13027).Bug fixes
jax.numpy.linalg.det and {func}jax.numpy.linalg.slogdet now use a
closed-form LU decomposition with row pivoting for 2x2 and 3x3 matrices
instead of closed-form polynomial expansions to avoid numerical instability
and catastrophic cancellation ({jax-issue}#39905).jax_explain_cache_misses is enabled and a
jax.custom_batching.custom_vmap-decorated function (or
custom_partitioning, custom_gradient, closure_convert,
linear_call, or run_state) is retraced with different-shaped
arguments ({jax-issue}#40110).jax.nn.dot_product_attention with implementation='cudnn') now
support operands that do not carry the vmap axis, including a shared
bias or mask. Previously jax.jacobian, jax.vmap with partial
in_axes, and jax.vmap of a VJP or of jax.grad failed with a
reshape TypeError ({jax-issue}#38495).jax.vmap of fp8 cuDNN fused attention now works: its batching rules
additionally mislabeled or dropped the amax outputs and restored output
shapes incorrectly, so previously no vmap of the fp8 path succeeded at
all. The amax outputs are whole-batch statistics and do not carry the
vmap axis; vmap over the scale/descale operands raises a clear
NotImplementedError.jax_compiler_enable_remat_pass to False now adds
rematerialization to the set of disabled XLA passes instead of
overwriting it, so HLO passes disabled via
XLA_FLAGS=--xla_disable_hlo_passes=... stay disabled
({jax-issue}#37391).jax.numpy.split, {func}jax.numpy.array_split, and the
hsplit/vsplit/dsplit variants once again accept negative entries in
indices_or_sections, resolving them against the axis size as NumPy does
({jax-issue}#6599). Out-of-bound indices are now clipped to the axis
bounds and produce empty sections, also matching NumPy, instead of raising
ValueError: Sizes passed to split must be nonnegative.jax.lax.scan to only check .mat
equivalency when the abstract value is a ShapedArray
({jax-issue}#39700).jax.lax.reshape
when reshaping arrays with sharding constraints ({jax-issue}#39309).jax.tree_util.flatten_one_level_with_keys for namedtuple
instances ({jax-issue}#39297)._get_prime_factors in jax.experimental.mesh_utils
({jax-issue}#38286).PyTreeDef.deserialize_using_proto now raises ValueError for a
malformed PyTreeDefProto instead of crashing the interpreter. A node
whose arity exceeds the subtrees preceding it is rejected, as is a dict
node whose arity disagrees with its key list, which previously segfaulted
later, when the structure was used to unflatten.
PyTreeDef.compose likewise raises ValueError rather than crashing when
an operand carries such an inconsistent arity, which is reachable from
pickle.loads ({jax-issue}#37410).Deprecations
jax.experimental.pallas.ops.gpu are deprecated.
Please use tokamax for equivalent
implementations if available, e.g., tokamax.layer_norm,
tokamax.dot_product_attention.jax.experimental.pallas.core_map is deprecated. Please migrate to
{func}jax.experimental.pallas.kernel.New features
barrier argument of
{func}jax.experimental.pallas.mosaic_gpu.copy_gmem_to_smem must now be
omitted on pre-Hopper GPUs (which use the cp.async implementation).
When omitted, the completion of the copy must be awaited via
{func}jax.experimental.pallas.mosaic_gpu.wait_gmem_to_smem.The deprecated module j ax.cloud_tpu_init was removed. This did nothing and references to it can be safely removed.
New features
hijax-custom-derivatives), along withjax.experimental.hijax helpers for deriving VJPHiPrimitive autodiffjvp or lin rule: linearize_from_jvp withapply_derived_linearization, vjp_fwd_from_jvp with transpose_jvp,vjp_fwd_from_lin with transpose_linearized, and jvp_from_lin.jax.custom_remat to the top-level jax namespace, forjax_remat3jax.checkpoint_policies is now a submodule rather than a namespacefrom jax.checkpoint_policies import ... now works; attributeSaveOnlyTheseNames, SaveAnyNamesButThese, andSaveAndOffloadOnlyTheseNames.jax.Inline enum for specify inlining policies tojax.jit.Breaking changes
ax.cloud_tpu_init was removed. This did nothing and3.13t) has been dropped. Pythoncibuildwheel, scipy) are making similar moves.jax.numpy.empty and jax.numpy.empty_like now producejax.numpy.zeros or jax.numpy.zeros_like instead.Deprecations
jax.numpy.cross is deprecated and will be removed in JAX 0.12.0, aligning with NumPy 2.5 behavior.jax.core have been removed, includingCallPrimitive, DebugInfo, DropVar, Effect, Effects, InconclusiveDimensionOperation,JaxprTypeError, abstract_token, check_jaxpr, concrete_or_error, find_top_trace, gensym,get_opaque_trace_state, is_concrete, is_constant_dim, is_constant_shape, jaxprs_in_params,new_jaxpr_eqn, no_effects, nonempty_axis_env_DO_NOT_USE, primal_dtype_to_tangent_dtype,unsafe_am_i_under_a_jit_DO_NOT_USE, unsafe_am_i_under_a_vmap_DO_NOT_USE,unsafe_get_axis_names_DO_NOT_USE, valid_jaxtype, JaxprPpContext, JaxprPpSettings, OutputType,aval_mapping_handlers, call, concretization_function_error, custom_typechecks,literalable_types, no_axis_name, and trace_ctx.jax.interpreters.pxla have been removed, includingIndex, MeshAxisName, MeshExecutable, global_aval_to_result_handler, global_result_handlers,are_hlo_shardings_equal, is_hlo_sharding_replicated, ArrayMapping, _UNSPECIFIED,array_mapping_to_axis_resources, and op_sharding_to_indices.Changes
jax.experimental.pallas.kernel now only aliases closed over or
passed in refs if the kernel writes to them. If you need aliasing without
an explicit write, use {func}jax.experimental.pallas.tpu.touch.The Triton backend is deprecated and will be removed in a future version of JAX.
To keep using Pallas on GPU, please migrate to the Mosaic GPU backend. To keep
using Triton, switch to the official Triton bindings and jax_triton.
New features
jax.experimental.pallas.triton.CUSTOM_CALL_TARGET_NAME, the
name of the custom call targeted by Pallas Triton. The exact value is not
guaranteed to be stable, but it is safe to assume that it is always aligned
with what the lowering emits. This is useful, for example, when exporting
a computation with Pallas Triton kernels via
{func}jax.experimental.export.export, since the corresponding custom call
is not considered export-stable and needs to be enabled explicitly.New features
jax.experimental.pallas.mosaic_gpu.mma.cp.async to
{func}jax.experimental.pallas.mosaic_gpu.copy_gmem_to_smem.jax.experimental.pallas.mosaic_gpu.wgmma now supports mixing the
e4m3 and e5m2 FP8 types across its two operands.jax.experimental.pallas.mosaic_gpu.emit_pipeline_warp_specialized.oob_fill_mode to
{class}jax.experimental.pallas.mosaic_gpu.BlockSpec, which controls
the behavior for out-of-bounds accesses during pipelined GMEM-to-SMEM
copies.Changes
jax.experimental.pallas.mosaic_gpu.LoweringSemantics.Lane in your
kernel's {class}jax.experimental.pallas.mosaic_gpu.CompilerParams and file
a bug.Deprecations
idx parameter of {func}jax.experimental.pallas.mosaic_gpu.load is
deprecated. Index the ref explicitly via ref.at[idx] prior to loading
from it.Added jax.scipy.linalg.invhilbert for the closed-form inverse of the Hilbert matrix ({jax-issue} #10144 ).
jax.scipy.linalg.invhilbert for the closed-form inverse#10144).jax.scipy.linalg.invpascal for the inverse of the Pascal#10144).jax.scipy.linalg.fiedler_companion for constructing the#10144).jax.ShapeDtypeStruct.like -- a shortcut for constructing ajax.ShapeDtypeStruct from an object with shape and dtypejax.scipy.linalg.invhilbert for the closed-form inverse
of the Hilbert matrix ({jax-issue}#10144).jax.scipy.linalg.invpascal for the inverse of the Pascal
matrix ({jax-issue}#10144).jax.scipy.linalg.fiedler_companion for constructing the
pentadiagonal Fiedler companion matrix of a polynomial
({jax-issue}#10144).jax.ShapeDtypeStruct.like -- a shortcut for constructing a
{class}jax.ShapeDtypeStruct from an object with shape and dtype
attributes.New features
jax.experimental.pallas.enable_poison_buffers (config flag
jax_pallas_poison_buffers) to poison (initialize with NaNs or lowest
possible integers) any scratch buffers allocated by Pallas for debugging.Deprecations
pl.debug_checks_enabled is deprecated. Use pl.enable_debug_checks.value.pl.dot was moved into {mod}jax.experimental.pallas.triton.
Accessing it via {mod}jax.experimental.pallas is deprecated.
You can use {func}jax.numpy.dot, {func}jax.numpy.einsum or the @
operator instead in a TPU or MGPU kernel.New features
jax_pallas_auto_assign_collective_ids config flag to allow two new
custom semaphore barrier collective IDs modes: ('yes') assigning missing
collective IDs automatically or ('override') overridding all collective IDs
and assigning them automatically, both based on the serialized kernel hash.Changes
jax.experimental.pallas.tpu.CompilerParams now defaults
needs_layout_passes to True. The layout passes are still a work in
progress. Please file a bug if you encounter a compilation error with
them enabled.jax.experimental.pallas.tpu.CompilerParams now defaults
use_tc_tiling_on_sc to True for SparseCore kernels.Removals
jax.experimental.pallas.tpu.KernelType and
{func}jax.experimental.pallas.tpu.repeat.Deprecations
pltpu.HOST and pltpu.MemorySpace.HOST in favor of pl.HOST.New features
jax.experimental.pallas.multiple_of to specify
divisibility requirements on dynamic indices.plgpu.Barriers and
plgpu.ClusterBarriers, by providing a nD shape as the num_barriers
parameter.plgpu.async_store_tmem support for TMEM
references with batch dimensions.TMA_GATHER_INDICES_LAYOUT to TMA_INDICES_LAYOUT.Changes
jax.experimental.pallas.program_id and
{func}jax.experimental.pallas.num_programs no longer work inside
kernels defined via {func}jax.experimental.pallas.mosaic_gpu.kernel.
Use {func}jax.lax.axis_index and {func}jax.lax.axis_size instead.jax.experimental.pallas.mosaic_gpu.kernel now has the same API as
{func}jax.experimental.pallas.kernel, meaning that it can be used as
a decorator and also uses out_type and scratch_types instead of
out_shape and scratch_shapes.Removals
plgpu.unswizzle_ref and plgpu.untile_ref.Deprecations
jax.experimental.pallas.pallas_call for Mosaic GPU kernels
is deprecated. Please migrate to
{func}jax.experimental.pallas.mosaic_gpu.kernel and
{func}jax.experimental.pallas.mosaic_gpu.emit_pipeline.with mesh: context manager has been deprecated. Please use with jax.set_mesh(mesh): instead.
New features
ResizeMethod.AREA to jax.image.resize, which matchesjax.scipy.linalg.hadamard for constructing Hadamardjax.scipy.linalg.circulant for constructing circulantjax.scipy.linalg.dft for constructing discrete Fourierjax.scipy.linalg.leslie for constructing Leslie matricesjax.scipy.linalg.companion for constructing companionjax.scipy.linalg.fiedler for constructing symmetric Fiedlerjax.scipy.linalg.helmert for constructing Helmert matricesjax.random.key_dtype to get the dtype corresponding to a PRNGjax.random.key and wrap_key_data now accept a dtype argument.Breaking changes
with mesh: context manager has been deprecated. Please usewith jax.set_mesh(mesh): instead.Deprecations
copy, order, and ndmin arguments tojax.numpy.array positionally is deprecated. Use keyword argumentsnumpy.array.dict_values, generators, zip return type and iterators generallyis_leaf to jax.tree.* methods.New features
ResizeMethod.AREA to {func}jax.image.resize, which matches
TensorFlow's AREA resizing ({jax-issue}#20098).jax.scipy.linalg.hadamard for constructing Hadamard
matrices ({jax-issue}#10144).jax.scipy.linalg.circulant for constructing circulant
matrices ({jax-issue}#10144).jax.scipy.linalg.dft for constructing discrete Fourier
transform matrices ({jax-issue}#10144).jax.scipy.linalg.leslie for constructing Leslie matrices
({jax-issue}#10144).jax.scipy.linalg.companion for constructing companion
matrices from polynomial coefficients ({jax-issue}#10144).jax.scipy.linalg.fiedler for constructing symmetric Fiedler
matrices ({jax-issue}#10144).jax.scipy.linalg.helmert for constructing Helmert matrices
({jax-issue}#10144).jax.scipy.special.boxcox and
{func}jax.scipy.special.boxcox1p for the Box-Cox power transformation.#27854):
jax.random.key_dtype to get the dtype corresponding to a PRNG
implementation name.jax.random.key and wrap_key_data now accept a dtype argument.Breaking changes
with mesh: context manager has been deprecated. Please use
with jax.set_mesh(mesh): instead.Deprecations
copy, order, and ndmin arguments to
{func}jax.numpy.array positionally is deprecated. Use keyword arguments
instead. This matches the signature of numpy.array.dict_values, generators, zip return type and iterators generally
are deprecated by default when used as leaves in pytrees. In a future
version of JAX, this will become an error, if you depend on using them as
leaves, pass is_leaf to jax.tree.* methods.Changes
jax.experimental.pallas.align_to, a utility that rounds a
value up to the nearest multiple of a given alignment.jax.experimental.pallas.pallas_call no longer supports checkify.
We expect this change to affect few users, as in our experience most
kernels either perform no checking or use
{func}jax.experimental.pallas.debug_check for conditionally-enabled
runtime checks.jax.experimental.pallas.kernel now always aliases Refs that are
passed in or closed-over.Removals
kernel_type field from
{class}jax.experimental.pallas.tpu.CompilerParams. It was only used for
writing SparseCore kernels via {func}jax.experimental.pallas.pallas_call,
which is now unsupported. The recommended API for SparseCore kernels is
{func}jax.experimental.pallas.kernel.New features
jax.experimental.pallas.mosaic_gpu.barrier_test function; a
non-blocking equivalent of
{func}jax.experimental.pallas.mosaic_gpu.barrier_wait only supported in
a warp context.Changes
plgpu.TransposeTransform.The deprecated keyword arguments a , a_min , and a_max to jax.numpy.clip have been removed.
New features:
ResizeMethod.CUBIC_PYTORCH to jax.image.resize to matchfull_matrices is True.perturb_singular argument toBreaking changes:
PartitionSpec objects no longer report themselves to be equal to tuples.PartitionSpec objects before testing equality..vma property has been removed from jax.core.ShapedArray. Use.manual_axis_type.varying instead.cpu:0, cpu:1, etc. instead ofTFRT_CPU_0, TFRT_CPU_1.jax_pmap_shmap_merge has been removed. jax.pmapjax.jit(jax.shard_map). Please seejax.device_put_sharded and jax.device_put_replicated have been removedAttributeError when accessed.jax.sharding.PmapShardingjaxlib.xla_extension: PmapFunction, pmap,NoSharding, Chunked, Unstacked, ShardedAxis, Replicated,ShardingSpec.jax.interpreters.pxla: MapTracer, PmapExecutable,parallel_callable, shard_args, xla_pmap_p, Chunked,NoSharding, Replicated, ShardedAxis, ShardingSpec,Unstacked, spec_to_indices.a, a_min, and a_max tojax.numpy.clip have been removed.jax.numpy.hstack, jax.numpy.vstack, jax.numpy.dstack,jax.numpy.column_stack, jax.numpy.atleast_1d, jax.numpy.atleast_2d,jax.numpy.atleast_3d no longer accept non-ArrayLike inputs.DeprecationWarning.Deprecations:
jax.core have been newly deprecated andjax.extend.core. These include CallPrimitive,DebugInfo, DropVar, Effect, Effects, InconclusiveDimensionOperation,JaxprTypeError, check_jaxpr, concrete_or_error, find_top_trace,gensym, get_opaque_trace_state, jaxprs_in_params, new_jaxpr_eqn,no_effects, nonempty_axis_env_DO_NOT_USE, primal_dtype_to_tangent_dtype,unsafe_am_i_under_a_jit_DO_NOT_USE, unsafe_am_i_under_a_vmap_DO_NOT_USE,unsafe_get_axis_names_DO_NOT_USE, valid_jaxtype, JaxprPpContext,JaxprPpSettings, OutputType, abstract_token, aval_mapping_handlers,call, concretization_function_error, custom_typechecks, is_concrete,is_constant_dim, is_constant_shape, literalable_types, no_axis_name,pytype_aval_mappings, and trace_ctx.Changes:
vma parameter of jax.ShapeDtypeStruct has been replaced withmanual_axis_type: jax.sharding.ManualAxisType. The .vma property has.manual_axis_type.varying.jax.scipy.linalg.cho_solve, jax.scipy.linalg.lu_solve, andjax.scipy.linalg.solve_triangular now show a deprecation warning forb.ndim > 1. In the future these will be treated asBug fixes:
jax.lax.linalg.tridiagonal_solve on GPU (#32487).jax.scipy.fft.dctn and idctn where axes=Nones was specified, instead of thelen(s) axes to match SciPy behavior (#29426).jax.distributed.initialize() on a GCE TPUIndexError (#36593). Whenjax.distributed.initialize() is called on a GCE VM, it uses the GCEChanges
pl.kernel to use out_type instead of
out_shape and scratch_types instead of scratch_shapes. Existing
usages calling {func}jax.experimental.pallas.kernel with these keyword
arguments must be updated.Deprecations
pltpu.semaphore, pltpu.DeviceIdType, pltpu.semaphore_signal,
pltpu.semaphore_wait, and pltpu.semaphore_read are now available in
{mod}jax.experimental.pallas. Accessing them via
{mod}jax.experimental.pallas.tpu is deprecated.Removals
pltpu.ANY and pltpu.MemorySpace.ANY.
Use pl.ANY instead.pltpu.delay, which is now available as pl.delay.New features
leader_tracked argument to ClusterBarrier, which allows
tracking barrier completions solely from the leader block along a specific
axis in a cluster.The semi-private type jax._src.literals.TypedNdArray is now a subclass of np.ndarray , rather than a duck type of it.
jax._src.literals.TypedNdArray is now a subclass ofnp.ndarray, rather than a duck type of it.jax.numpy.arange with step specified no longer generates the arrayjnp.array(np.arange(...)).New features:
Added support for atomics to the Mosaic GPU backend, through
{func}jax.experimental.pallas.mosaic_gpu.atomic_add,
{func}jax.experimental.pallas.mosaic_gpu.atomic_max,
{func}jax.experimental.pallas.mosaic_gpu.atomic_min,
{func}jax.experimental.pallas.mosaic_gpu.atomic_and,
{func}jax.experimental.pallas.mosaic_gpu.atomic_or, and
{func}jax.experimental.pallas.mosaic_gpu.atomic_xor.
Added a leader_tracked argument to
{func}jax.experimental.pallas.mosaic_gpu.copy_gmem_to_smem. This adds
support for cluster-collective partitioned/replicated GMEM to SMEM copies
whose completion is tracked only by the leader block. Removed the
partitioned_axis argument (which provided limited access to this feature)
from the API.
JAX tracers that are not of Array type (e.g., of Ref type) will no longer report themselves to be instances of Array .
Changes:
Array type (e.g., of Ref type) will noArray.jax.shard_map in Explicit mode will raise an errorin_specs. In other words, it will act like an assert instead of anin_specs is an optional argument so you can omit specifying itshard_map will infer the PartitionSpec from the argument. If youjax.reshard on the arguments andNew features:
jax_compilation_cache_check_contents. If set, we missget() is called on a value that has not been put() by the currentput(), we verify that its contents match.New features:
jax.experimental.pallas.with_scoped decorator that provides
the function with scope-allocated scratch buffers.Changes
backend argument of
{func}jax.experimental.pallas.pallas_call in favor of compiler_params.
For example, to force the use of the Triton backend you have to now write
compiler_params=pltriton.CompilerParams(), where pltriton refers to
{mod}jax.experimental.pallas.triton.jax.experimental.pallas.tpu.KernelType to CoreType. The
old name is deprecated and will be removed in a future release.JAX v0.9.0.1 is identical to v0.9.0 with the commits from the following four PRs patched in:
JAX v0.9.0.1 is identical to v0.9.0 with the commits from the following four PRs patched in:
Setting the jax_pmap_shmap_merge config state is deprecated in JAX v0.9.0 and will be removed in JAX v0.10.0.
New features:
jax.thread_guard, a context manager that detects when devicesBug fixes:
magma_zgeqp3_gpu)use_magma=True and pivoting=True.Deprecations:
jax_collectives_common_channel_id was removed.jax_pmap_no_rank_reduction config state has been removed. Thejax.pmapped function f sees inputs of the same rank as the input tojax.pmap(f). For example, if jax.pmap(f) receives shape (8, 128) onf receives shape (1, 128).jax_pmap_shmap_merge config state is deprecated in JAX v0.9.0jax.numpy.fix is deprecated, anticipating the deprecation ofnumpy.fix in NumPy v2.5.0. jax.numpy.trunc is a drop-inChanges:
jax.export now supports explicit sharding. This required a newNew features:
reduction_scratch_bytes field to
{class}jax.experimental.pallas.mosaic_gpu.CompilerParams. This gives user
control over how much shared memory Pallas is allowed to reserve for
cross-warp reductions on GPU. Increasing this value typically allows for
faster reductions.Changes
jax.experimental.pallas.pallas_call with
the backend argument set to 'triton'.Removals
pl.atomic_*, pl.load, pl.store,
pl.swap and pl.max_contiguous.JAX v0.8.3 is identical to v0.8.2 with the following two bug fixes patched in:
JAX v0.8.3 is identical to v0.8.2 with the following two bug fixes patched in:
jax.lax.pvary has been deprecated. Please use jax.lax.pcast(..., to='varying') as the replacement.
Deprecations
jax.lax.pvary has been deprecated.
Please use jax.lax.pcast(..., to='varying') as the replacement.jax.numpy.arange now result in a
deprecation warning, because the output is poorly-defined.jax.core a number of symbols are newly deprecated including:
call_impl, get_aval, mapped_aval, subjaxprs, set_current_trace,
take_current_trace, traverse_jaxpr_params, unmapped_aval,
AbstractToken, and TraceTag.jax.interpreters.pxla are deprecated. These are
primarily JAX internal APIs, and users should not rely on them.Changes:
jax's Tracer no longer inherits from jax.Array at runtime. However,
jax.Array now uses a custom metaclass such isinstance(x, Array) is true
if an object x represents a traced Array. Only some Tracers represent
Arrays, so it is not correct for Tracer to inherit from Array.
For the moment, during Python type checking, we continue to declare Tracer
as a subclass of Array, however we expect to remove this in a future
release.
jax.experimental.si_vjp has been deleted.
jax.vjp subsumes it's functionality.
jax.sharding.PmapSharding is now deprecated. Please use jax.NamedSharding instead.
New features:
jax.jit now supports the decorator factory pattern; i.e instead of
writing@functools.partial(jax.jit, static_argnames=['n'])
def f(x, n):
...
you may write@jax.jit(static_argnames=['n'])
def f(x, n):
...
Changes:
{func}jax.lax.linalg.eigh now accepts an implementation argument to
select between QR (CPU/GPU), Jacobi (GPU/TPU), and QDWH (TPU)
implementations. The EighImplementation enum is publicly exported from
{mod}jax.lax.linalg.
{func}jax.lax.linalg.svd now implements an algorithm that uses the polar
decomposition on CUDA GPUs. This is also an alias for the existing algorithm
on TPUs.
Bug fixes:
#33062).Deprecations:
jax.sharding.PmapSharding is now deprecated. Please use
jax.NamedSharding instead.jx.device_put_replicated is now deprecated. Please use jax.device_put
with the appropriate sharding instead.jax.device_put_sharded is now deprecated. Please use jax.device_put with
the appropriate sharding instead.axis_types of jax.make_mesh will change in JAX v0.9.0 to return
jax.sharding.AxisType.Explicit. Leaving axis_types unspecified will raise a
DeprecationWarning.jax.cloud_tpu_init and its contents were deprecated. There is no reason for a user to import or use the contents of this module; JAX handles this for you automatically if needed.New features:
jax.experimental.pallas.tpu.get_tpu_info to get TPU hardware information.Deprecations
pl.max_contiguous has been moved to {mod}jax.experimental.pallas.triton.
Accessing it via {mod}jax.experimental.pallas is deprecated.pl.swap is deprecated and will be removed in a future release. Use
indexing or backend-specific loading/storing APIs instead.Removals
jax.experimental.pallas.tpu.TPUCompilerParams,
{class}jax.experimental.pallas.tpu.TPUMemorySpace,
{class}jax.experimental.pallas.tpu.TritonCompilerParams.The deprecated function {func}jax.interpreters.mlir.custom_call was removed.
Breaking changes:
jax.pmap implementation to one implemented in
terms of jax.jit and jax.shard_map. jax.pmap is in maintenance mode
and we encourage all new code to use jax.shard_map directly. See the
migration guide for
more information.auto= parameter of jax.experimental.shard_map.shard_map has been
removed. This means that jax.experimental.shard_map.shard_map no longer
supports nesting. If you want to nest shard_map calls, please use
jax.shard_map.__jax_array__ directly
to, e.g. jit-ed functions. Call jax.numpy.asarray on them first.jax.numpy.cov is now returns NaN for empty arrays ({jax-issue}#32305),
and matches NumPy 2.2 behavior for single-row design matrices ({jax-issue}#32308).Array values where a dtype value is expected. Call
.dtype on these values first.jax.interpreters.mlir.custom_call was
removed.jax.util, jax.extend.ffi, and jax.experimental.host_callback
modules have been removed. All public APIs within these modules were
deprecated and removed in v0.7.0 or earlier.jax.custom_derivatives.custom_jvp_call_jaxpr_p
was removed.jax.experimental.multihost_utils.process_allgather raises an error when
the input is a jax.Array and not fully-addressable and tiled=False. To fix
this, pass tiled=True to your process_allgather invocation.jax.experimental.compilation_cache, the deprecated symbols
is_initialized and initialize_cache were removed.jax.interpreters.xla.canonicalize_dtype
was removed.jaxlib.hlo_helpers has been removed. Use {mod}jax.ffi instead.jax_cpu_enable_gloo_collectives has been removed. Use
jax_cpu_collectives_implementation instead.interpolation argument to
{func}jax.numpy.percentile and {func}jax.numpy.quantile has been
removed; use method instead.for_loop primitive was removed. Its functionality,
reading from and writing to refs in the loop body, is now directly
supported by {func}jax.lax.fori_loop. If you need help updating your
code, please file a bug.jax.numpy.trimzeros now errors for non-1D input.where argument to {func}jax.numpy.sum and other reductions is now
required to be boolean. Non-boolean values have resulted in a
DeprecationWarning since JAX v0.5.0.jax.dlpack, {mod} jax.errors, {mod}
jax.lib.xla_bridge, {mod} jax.lib.xla_client, and {mod}
jax.lib.xla_extension were removed.jax.interpreters.mlir.dense_bool_array was removed. Use MLIR APIs to
construct attributes instead.Changes
jax.numpy.linalg.eig now returns a namedtuple (with attributes
eigenvalues and eigenvectors) instead of a plain tuple.jax.grad and {func}jax.vjp will now round always primals to
float32 if float64 mode is not enabled.jax.dlpack.from_dlpack now accepts arrays with non-default layouts,
for example, transposed.implementation argument to {func}jax.lax.linalg.eig
({jax-issue}#27265). The use_magma argument is now deprecated in favor
of implementation.jax.numpy.trim_zeros now follows NumPy 2.2 in supporting
multi-dimensional inputs.Deprecations
jax.experimental.enable_x64 and {func}jax.experimental.disable_x64
are deprecated in favor of the new non-experimental context manager
{func}jax.enable_x64.jax.experimental.shard_map.shard_map is deprecated; going forward use
{func}jax.shard_map.jax.experimental.pjit.pjit is deprecated; going forward use
{func}jax.jit.{func}jax.dlpack.from_dlpack no longer accepts a DLPack capsule. This behavior was deprecated and is now removed. The function must be called with an…
Breaking changes:
jax.dlpack.from_dlpack no longer accepts a DLPack capsule. This
behavior was deprecated and is now removed. The function must be called
with an array implementing __dlpack__ and __dlpack_device__.Changes
The minimum supported NumPy version is now 2.0. Since SciPy 1.13 is required for NumPy 2.0 support, the minimum supported SciPy version is now 1.13.
JAX now represents constants in its internal jaxpr representation as a
TypedNdArray, which is a private JAX type that duck types as a
numpy.ndarray. This type may be exposed to users via custom_jvp rules,
for example, and may break code that uses isinstance(x, np.ndarray). If
this breaks your code, you may convert these arrays to classic NumPy arrays
using np.asarray(x).
Bug fixes
arr.view(dtype=None) now returns the array unchanged, matching NumPy's
semantics. Previously it returned the array with a float dtype.jax.random.randint now produces a less-biased distribution for 8-bit and
16-bit integer types ({jax-issue}#27742). To restore the previous biased
behavior, you may temporarily set the jax_safer_randint configuration to
False, but note this is a temporary config that will be removed in a
future release.Deprecations:
enable_xla and native_serialization for jax2tf.convert
are deprecated and will be removed in a future version of JAX. These were
used for jax2tf with non-native serialization, which has been now removed.jax_pmap_no_rank_reduction to False is
deprecated. By default, jax_pmap_no_rank_reduction will be set to True
and jax.pmap shards will not have their rank reduced, keeping the same
rank as their enclosing array.{func}jax.lax.zeros_like_array is deprecated. Please use {func}jax.numpy.zeros_like instead.
New features
Changes
jax.set_mesh which acts as a global setter and a context manager.
Removed jax.sharding.use_mesh in favor of jax.set_mesh.jax.lax.dot now implements the general dot product via the optional
dimension_numbers argument.Deprecations:
jax.lax.zeros_like_array is deprecated. Please use
{func}jax.numpy.zeros_like instead.jax.experimental.host_callback now results in
a DeprecationWarning, and will result in an ImportError starting in JAX
v0.8.0. Its APIs have raised NotImplementedError since JAX version 0.4.35.jax.lax.dot, passing the precision and preferred_element_type
arguments by position is deprecated. Pass them by explicit keyword instead.jax.interpreters.ad,
{mod}jax.interpreters.batching, and {mod}jax.interpreters.partial_eval; they
are used rarely if ever outside JAX itself, and most are deprecated without any
public replacement.New features:
pltpu.make_async_remote_copy and pltpu.semaphore_signal's device_id
argument now allows user to pass in a dictionary that only specifies the
device index along the communication axis, instead of the full coordinates.
It also supports TPU core id index.jax.debug.print now works in Pallas kernels and is the recommended way to
print.Deprecations
pl.atomic_* APIs have been moved to {mod}jax.experimental.pallas.triton.
Accessing them via {mod}jax.experimental.pallas is deprecated.pl.load and pl.store are deprecated. Use indexing or backend-specific
loading/storing APIs instead.Doing otherwise will result in an error starting in v0.7.x. This raised a DeprecationWarning in v0.6.x.
New features:
jax.P which is an alias for jax.sharding.PartitionSpec.jax.tree.reduce_associative.jax.numpy.ndarray.at indexing methods now support a wrap_negative_indices
argument, which defaults to True to match the current behavior ({jax-issue}#29434).Breaking changes:
jax.stages.OutInfo has been replaced with jax.ShapeDtypeStruct.jax.jit now requires fun to be passed by position, and additional
arguments to be passed by keyword. Doing otherwise will result in an error
starting in v0.7.x. This raised a DeprecationWarning in v0.6.x.Layout, .layout, .input_layouts and .output_layouts have been
renamed to Format, .format, .input_formats and .output_formatsDeviceLocalLayout, .device_local_layout have been renamed to Layout
and .layoutjax.experimental.shard module has been deleted and all the APIs have been
moved to the jax.sharding endpoint. So use jax.sharding.reshard,
jax.sharding.auto_axes and jax.sharding.explicit_axes instead of their
experimental endpoints.lax.infeed and lax.outfeed were removed, after being deprecated in
JAX 0.6. The transfer_to_infeed and transfer_from_outfeed methods were
also removed the Device objects.jax.extend.core.primitives.pjit_p primitive has been renamed to
jit_p, and its name attribute has changed from "pjit" to "jit".
This affects the string representations of jaxprs. The same primitive is no
longer exported from the jax.experimental.pjit module.jax.extend.backend.add_clear_backends_callback
has been removed. Users should use jax.extend.backend.register_backend_cache
instead.out_sharding arg added to x.at[y].set and x.at[y].add. Previous
behavior propagating operand sharding removed. Please use
x.at[y].set/add(z, out_sharding=jax.typeof(x).sharding) to retain previous
behavior if scatter op requires collectives.Deprecations:
jax.dlpack.SUPPORTED_DTYPES is deprecated; please use the new
{func}jax.dlpack.is_supported_dtype function.jax.scipy.special.sph_harm has been deprecated following a similar
deprecation in SciPy; use {func}jax.scipy.special.sph_harm_y instead.jax.interpreters.xla, the previously deprecated symbols
abstractify and pytype_aval_mappings have been removed.jax.interpreters.xla.canonicalize_dtype is deprecated. For
canonicalizing dtypes, prefer {func}jax.dtypes.canonicalize_dtype.
For checking whether an object is a valid jax input, prefer
{func}jax.core.valid_jaxtype.jax.core, the previously deprecated symbols AxisName,
ConcretizationTypeError, axis_frame, call_p, closed_call_p,
get_type, trace_state_clean, typematch, and typecheck have been
removed.jax.lib.xla_client, the previously deprecated symbols
DeviceAssignment, get_topology_for_devices, and mlir_api_version
have been removed.jax.extend.ffi was removed after being deprecated in v0.5.0.
Use {mod}jax.ffi instead.jax.lib.xla_bridge.get_compile_options is deprecated, and replaced by
{func}jax.extend.backend.get_compile_options.New functionality
jax.experimental.pallas.loop which allows
to write stateless loops as functions.jax.experimental.pallas.tpu.emit_pipeline. Input buffers can now
be multiple-buffered with more than 2 buffers and support a lookahead option
to fetch blocks that are an arbitrary number of grid iterations ahead
rather than the immediate next iterations. Additionally, pipeline state
can now be held in registers to reduce scalar memory usage.Deprecations
jax.experimental.pallas.triton.TritonCompilerParams has been
renamed to {class}jax.experimental.pallas.triton.CompilerParams. The
old name is deprecated and will be removed in a future release.jax.experimental.pallas.tpu.TPUCompilerParams
and {class}jax.experimental.pallas.tpu.TPUMemorySpace have been
renamed to {class}jax.experimental.pallas.tpu.CompilerParams
and {class}jax.experimental.pallas.tpu.MemorySpace. The
old names are deprecated and will be removed in a future release.Added {func}jax.tree.broadcast which implements a pytree prefix broadcasting helper.
New features:
jax.tree.broadcast which implements a pytree prefix broadcasting helper.Changes
jax.custom_derivatives.custom_jvp_call_jaxpr_p is deprecated, and will be removed in JAX v0.7.0.
New features:
jax.lax.axis_size which returns the size of the mapped axis
given its name.Changes
jax.sharding.PartitionSpec no longer inherits from a tuple.jax.ShapeDtypeStruct is immutable now. Please use .update method to
update your ShapeDtypeStruct instead of doing in-place updates.Deprecations
jax.custom_derivatives.custom_jvp_call_jaxpr_p is deprecated, and will be
removed in JAX v0.7.0.Removals
jax.experimental.pallas.gpu. To use
the Triton backend import {mod}jax.experimental.pallas.triton.Changes
jax.experimental.pallas.BlockSpec now takes in special types in
addition to ints/None in the block_shape. indexing_mode has been
removed. To achieve "Unblocked", pass a pl.Element(size) into
block_shape for each entry that needs unblocked indexing.jax.experimental.pallas.pallas_call now requires compiler_params
to be a backend-specific dataclass instead of a param to value mapping.jax.experimental.pallas.debug_check is now supported both on
TPU and Mosaic GPU. Previously, this functionality was only supported
on TPU and required using the APIs from {mod}jax.experimental.checkify.
Note that debug checks are not executed unless
{data}jax.experimental.pallas.enable_debug_checks is set.{func}jax.numpy.array no longer accepts None. This behavior was deprecated since November 2023 and is now removed.
Breaking changes
jax.numpy.array no longer accepts None. This behavior was
deprecated since November 2023 and is now removed.config.jax_data_dependent_tracing_fallback config option,
which was added temporarily in v0.4.36 to allow users to opt out of the
new "stackless" tracing machinery.config.jax_eager_pmap config option.lower and trace AOT APIs on the result
of jax.jit if there have been subsequent wrappers applied.
Previously this worked, but silently ignored the wrappers.
The workaround is to apply jax.jit last among the wrappers,
and similarly for jax.pmap.
See {jax-issue}#27873.cuda12_pip extra for jax has been removed; use pip install jax[cuda12]
instead.Changes
pip install jax[cuda12_local]
to install JAX, run pip install jax[cuda12-local] instead.jax.jit now requires fun to be passed by position, and additional
arguments to be passed by keyword. Doing otherwise will result in a
DeprecationWarning in v0.6.X, and an error in starting in v0.7.X.Deprecations
jax.tree_util.build_tree is deprecated. Use {func}jax.tree.unflatten
instead.jax.lib.xla_extension are now deprecated.jax.interpreters.mlir.hlo and jax.interpreters.mlir.func_dialect,
which were accidental exports, have been removed. If needed, they are
available from jax.extend.mlir.jax.interpreters.mlir.custom_call is deprecated. The APIs provided by
{mod}jax.ffi should be used instead.jax.ffi.ffi_call with inline arguments is no
longer supported. {func}~jax.ffi.ffi_call now unconditionally returns a
callable.jax.lib.xla_client are deprecated:
get_topology_for_devices, heap_profile, mlir_api_version, Client,
CompileOptions, DeviceAssignment, Frame, HloSharding, OpSharding,
Traceback.jax.util are deprecated:
HashableFunction, as_hashable_function, cache, safe_map, safe_zip,
split_dict, split_list, split_list_checked, split_merge, subvals,
toposort, unzip2, wrap_name, and wraps.jax.dlpack.to_dlpack has been deprecated. You can usually pass a JAX
Array directly to the from_dlpack function of another framework. If you
need the functionality of to_dlpack, use the __dlpack__ attribute of an
array.jax.lax.infeed, jax.lax.infeed_p, jax.lax.outfeed, and
jax.lax.outfeed_p are deprecated and will be removed in JAX v0.7.0.jax.lib.xla_client: ArrayImpl, FftType, PaddingType,
PrimitiveType, XlaBuilder, dtype_to_etype,
ops, register_custom_call_target, shape_from_pyval, Shape,
XlaComputation.jax.lib.xla_extension: ArrayImpl, XlaRuntimeError.jax: jax.treedef_is_leaf, jax.tree_flatten, jax.tree_map,
jax.tree_leaves, jax.tree_structure, jax.tree_transpose, and
jax.tree_unflatten. Replacements can be found in {mod}jax.tree or
{mod}jax.tree_util.jax.core: AxisSize, ClosedJaxpr, EvalTrace, InDBIdx, InputType,
Jaxpr, JaxprEqn, Literal, MapPrimitive, OpaqueTraceState, OutDBIdx,
Primitive, Token, TRACER_LEAK_DEBUGGER_WARNING, Var, concrete_aval,
dedup_referents, escaped_tracer_error, extend_axis_env_nd, full_lower, get_referent, jaxpr_as_fun, join_effects, lattice_join,
leaked_tracer_error, maybe_find_leaked_tracers, raise_to_shaped,
raise_to_shaped_mappings, reset_trace_state, str_eqn_compact,
substitute_vars_in_output_ty, typecompat, and used_axis_names_jaxpr. Most
have no public replacement, though a few are available at {mod}jax.extend.core.vectorized argument to {func}~jax.pure_callback and
{func}~jax.ffi.ffi_call. Use the vmap_method parameter instead.Added a allow_negative_indices option to {func}jax.lax.dynamic_slice, {func}jax.lax.dynamic_update_slice and related functions. The default is true, m
New Features
allow_negative_indices option to {func}jax.lax.dynamic_slice,
{func}jax.lax.dynamic_update_slice and related functions. The default is
true, matching the current behavior. If set to false, JAX does not need to
emit code clamping negative indices, which improves code size.replace option to {func}jax.random.categorical to enable sampling
without replacement.…core.DebugInfo kwarg. For a limited time, a DeprecationWarning is printed if jax.extend.linear_util.wrap_init is used without debugging info. A downst…
Breaking changes
@jax.jit
def f(x):
return x
# inp1.sharding is of type SingleDeviceSharding
inp1 = jnp.arange(8)
f(inp1)
mesh = jax.make_mesh((1,), ('x',))
# inp2.sharding is of type NamedSharding
inp2 = jax.device_put(jnp.arange(8), NamedSharding(mesh, P('x')))
f(inp2) # tracing cache miss
In the above example, calling f(inp1) and then f(inp2) will lead to a
tracing cache miss because the shardings have changed on the abstract values
while tracing.New Features
jax.experimental.custom_dce.custom_dce
decorator to support customizing the behavior of opaque functions under
JAX-level dead code elimination (DCE). See {jax-issue}#25956 for more
details.jax.lax: {func}jax.lax.reduce_sum,
{func}jax.lax.reduce_prod, {func}jax.lax.reduce_max, {func}jax.lax.reduce_min,
{func}jax.lax.reduce_and, {func}jax.lax.reduce_or, and {func}jax.lax.reduce_xor.jax.lax.linalg.qr, and {func}jax.scipy.linalg.qr, now support
column-pivoting on CPU and GPU. See {jax-issue}#20282 andjax.random.multinomial.
{jax-issue}#25955 for more details.Changes
JAX_CPU_COLLECTIVES_IMPLEMENTATION and JAX_NUM_CPU_DEVICES now work as
env vars. Before they could only be specified via jax.config or flags.JAX_CPU_COLLECTIVES_IMPLEMENTATION now defaults to 'gloo', meaning
multi-process CPU communication works out-of-the-box.jax[tpu] TPU extra no longer depends on the libtpu-nightly package.
This package may safely be removed if it is present on your machine; JAX now
uses libtpu instead.Deprecations
linear_util.wrap_init and the constructor
core.Jaxpr now must take a non-empty core.DebugInfo kwarg. For
a limited time, a DeprecationWarning is printed if
jax.extend.linear_util.wrap_init is used without debugging info.
A downstream effect of this several other internal functions need debug
info. This change does not affect public APIs.
See https://github.com/jax-ml/jax/issues/26480 for more detail.jax.numpy.ndim, {func}jax.numpy.shape, and {func}jax.numpy.size,
non-arraylike inputs (such as lists, tuples, etc.) are now deprecated.Bug fixes
sudo sh -c 'echo always > /sys/kernel/mm/transparent_hugepage/enabled').
We hope to improve this further in future releases.As of this release, JAX now uses effort-based versioning. Since this release makes a breaking change to PRNG key semantics that may require users to u…
As of this release, JAX now uses effort-based versioning. Since this release makes a breaking change to PRNG key semantics that may require users to update their code, we are bumping the "meso" version of JAX to signify this.
Breaking changes
Enable jax_threefry_partitionable by default (see
the update note).
This release drops support for Mac x86 wheels. Mac ARM of course remains supported. For a recent discussion, see https://github.com/jax-ml/jax/discussions/22936.
Two key factors motivated this decision:
We are open to re-adding support for Mac x86 if the community is willing to help support that platform: in particular, we would need the JAX test suite to pass cleanly on Mac x86 before we could ship releases again.
Changes:
jax.numpy.einsum now defaults to optimize='auto' rather than
optimize='optimal'. This avoids exponentially-scaling trace-time in
the case of many arguments ({jax-issue}#25214).jax.numpy.linalg.solve no longer supports batched 1D arguments
on the right hand side. To recover the previous behavior in these cases,
use solve(a, b[..., None]).squeeze(-1).New Features
jax.numpy.fft.fftn, {func}jax.numpy.fft.rfftn,
{func}jax.numpy.fft.ifftn, and {func}jax.numpy.fft.irfftn now support
transforms in more than 3 dimensions, which was previously the limit. See
{jax-issue}#25606 for more details.jax.ffi.register_ffi_type_id function..as_text() method now supports the debug_info option
to include debugging information, e.g., source location, in the output.Deprecations
jax.interpreters.xla, abstractify and pytype_aval_mappings
are now deprecated, having been replaced by symbols of the same name
in {mod}jax.core.jax.scipy.special.lpmn and {func}jax.scipy.special.lpmn_values
are deprecated, following their deprecation in SciPy v1.15.0. There are
no plans to replace these deprecated functions with new APIs.jax.extend.ffi submodule was moved to {mod}jax.ffi, and the
previous import path is deprecated.Deletions
jax_enable_memories flag has been deleted and the behavior of that flag
is on by default.jax.lib.xla_client, the previously-deprecated Device and
XlaRuntimeError symbols have been removed; instead use jax.Device
and jax.errors.JaxRuntimeError respectively.jax.experimental.array_api module has been removed after being
deprecated in JAX v0.4.32. Since that release, {mod}jax.numpy supports
the array API directly.New functionality
jax.experimental.pallas.debug_print on TPU.a number of APIs in the internal jax.core namespace have been deprecated. Most were no-ops, were little-used, or can be replaced by APIs of the same n…
Breaking Changes
XlaExecutable.cost_analysis now returns a dict[str, float] (instead of a
single-element list[dict[str, float]]).Changes:
jax.tree.flatten_with_path and jax.tree.map_with_path are added
as shortcuts of the corresponding tree_util functions.Deprecations
jax.core namespace have been deprecated.
Most were no-ops, were little-used, or can be replaced by APIs of the same
name in {mod}jax.extend.core; see the documentation for {mod}jax.extend
for information on the compatibility guarantees of these semi-public extensions.jax.core: check_eqn, check_type, check_valid_jaxtype, and
non_negative_dim.jax.lib.xla_bridge: xla_client and default_backend.jax.lib.xla_client: _xla and bfloat16.jax.numpy: round_.New Features
jax.export.export can be used for device-polymorphic export with
shardings constructed with {func}jax.sharding.AbstractMesh.
See the jax.export documentation.jax.lax.split. This is a primitive version of
{func}jax.numpy.split, added because it yields a more compact
transpose during automatic differentiation.Breaking Changes
XlaExecutable.cost_analysis now returns a dict[str, float] (instead of a single-element list[dict[str, float]] ).
Changes:
jax.tree.flatten_with_path and jax.tree.map_with_path are added as shortcuts of the corresponding tree_util functions.
Deprecations
a number of APIs in the internal jax.core namespace have been deprecated. Most were no-ops, were little-used, or can be replaced by APIs of the same name in jax.extend.core ; see the documentation for jax.extend for information on the compatibility guarantees of these semi-public extensions.
Several previously-deprecated APIs have been removed, including:
from jax.core : check_eqn , check_type , check_valid_jaxtype , and non_negative_dim .
from jax.lib.xla_bridge : xla_client and default_backend .
from jax.lib.xla_client : _xla and bfloat16 .
from jax.numpy : round_ .
New Features
jax.export.export() can be used for device-polymorphic export with shardings constructed with jax.sharding.AbstractMesh() . See the jax.export documentation .
Added jax.lax.split() . This is a primitive version of jax.numpy.split() , added because it yields a more compact transpose during automatic differentiation.
{func}jax.experimental.jax2tf.convert with native_serialization=False or with enable_xla=False have been deprecated since July 2024, with JAX version…
Breaking Changes
This release lands "stackless", an internal change to JAX's tracing
machinery. We made trace dispatch purely a function of context rather than a
function of both context and data. This let us delete a lot of machinery for
managing data-dependent tracing: levels, sublevels, post_process_call,
new_base_main, custom_bind, and so on. The change should only affect
users that use JAX internals.
If you do use JAX internals then you may need to
update your code (see
https://github.com/jax-ml/jax/commit/c36e1f7c1ad4782060cbc8e8c596d85dfb83986f
for clues about how to do this). There might also be version skew
issues with JAX libraries that do this. If you find this change breaks your
non-JAX-internals-using code then try the
config.jax_data_dependent_tracing_fallback flag as a workaround, and if
you need help updating your code then please file a bug.
{func}jax.experimental.jax2tf.convert with native_serialization=False
or with enable_xla=False have been deprecated since July 2024, with
JAX version 0.4.31. Now we removed support for these use cases. jax2tf
with native serialization will still be supported.
In jax.interpreters.xla, the xb, xc, and xe symbols have been removed
after being deprecated in JAX v0.4.31. Instead use xb = jax.lib.xla_bridge,
xc = jax.lib.xla_client, and xe = jax.lib.xla_extension.
The deprecated module jax.experimental.export has been removed. It was replaced
by {mod}jax.export in JAX v0.4.30. See the migration guide
for information on migrating to the new API.
The initial argument to {func}jax.nn.softmax and {func}jax.nn.log_softmax
has been removed, after being deprecated in v0.4.27.
Calling np.asarray on typed PRNG keys (i.e. keys produced by {func}jax.random.key)
now raises an error. Previously, this returned a scalar object array.
The following deprecated methods and functions in {mod}jax.export have
been removed:
jax.export.DisabledSafetyCheck.shape_assertions: it had no effect
already.jax.export.Exported.lowering_platforms: use platforms.jax.export.Exported.mlir_module_serialization_version:
use calling_convention_version.jax.export.Exported.uses_shape_polymorphism:
use uses_global_constants.lowering_platforms kwarg for {func}jax.export.export: use
platforms instead.The kwargs symbolic_scope and symbolic_constraints from
{func}jax.export.symbolic_args_specs have been removed. They were
deprecated in June 2024. Use scope and constraints instead.
Hashing of tracers, which has been deprecated since version 0.4.30, now
results in a TypeError.
Refactor: JAX build CLI (build/build.py) now uses a subcommand structure and
replaces previous build.py usage. Run python build/build.py --help for
more details. Brief overview of the new subcommand options:
build: Builds JAX wheel packages. For e.g., python build/build.py build --wheels=jaxlib,jax-cuda-pjrtrequirements_update: Updates requirements_lock.txt files.{func}jax.scipy.linalg.toeplitz now does implicit batching on multi-dimensional
inputs. To recover the previous behavior, you can call {func}jax.numpy.ravel
on the function inputs.
{func}jax.scipy.special.gamma and {func}jax.scipy.special.gammasgn now
return NaN for negative integer inputs, to match the behavior of SciPy from
https://github.com/scipy/scipy/pull/21827.
jax.clear_backends was removed after being deprecated in v0.4.26.
We removed the custom call "__gpu$xla.gpu.triton" from the list of custom
call that we guarantee export stability. This is because this custom call
relies on Triton IR, which is not guaranteed to be stable. If you need
to export code that uses this custom call, you can use the disabled_checks
parameter. See more details in the documentation.
New Features
jax.jit got a new compiler_options: dict[str, Any] argument, for
passing compilation options to XLA. For the moment it's undocumented and
may be in flux.jax.tree_util.register_dataclass now allows metadata fields to be
declared inline via {func}dataclasses.field. See the function documentation
for examples.jax.numpy.put_along_axis.jax.lax.linalg.eig and the related jax.numpy functions
({func}jax.numpy.linalg.eig and {func}jax.numpy.linalg.eigvals) are now
supported on GPU. See {jax-issue}#24663 for more details.jax_exec_time_optimization_effort and jax_memory_fitting_effort, to control the amount of effort the compiler spends minimizing execution time and memory usage, respectively. Valid values are between -1.0 and 1.0, default is 0.0.jax.scipy.special.loggammaBug fixes
#24843 for more details.Deprecations
jax.lib.xla_extension.ArrayImpl and jax.lib.xla_client.ArrayImpl are deprecated;
use jax.Array instead.jax.lib.xla_extension.XlaRuntimeError is deprecated; use jax.errors.JaxRuntimeError
instead.jax.experimental.host_callback has been deprecated since March 2024, with JAX version 0.4.26. Now we removed it. See {jax-issue}#20385 for a discussio…
Breaking Changes
jax.numpy.isscalar now returns True for any array-like object with
zero dimensions. Previously it only returned True for zero-dimensional
array-like objects with a weak dtype.jax.experimental.host_callback has been deprecated since March 2024, with
JAX version 0.4.26. Now we removed it.
See {jax-issue}#20385 for a discussion of alternatives.Changes:
jax.lax.FftType was introduced as a public name for the enum of FFT
operations. The semi-public API jax.lib.xla_client.FftType has been
deprecated.libtpu package rather than
libtpu-nightly. For the next few releases JAX will pin an empty version of
libtpu-nightly as well as libtpu to ease the transition; that dependency
will be removed in Q1 2025.Deprecations:
jax.lib.xla_client.PaddingType has been deprecated.
No JAX APIs consume this type, so there is no replacement.jax.pure_callback and
{func}jax.extend.ffi.ffi_call under vmap has been deprecated and so has
the vectorized parameter to those functions. The vmap_method parameter
should be used instead for better defined behavior. See the discussion in
{jax-issue}#23881 for more details.jax.lib.xla_client.register_custom_call_target has
been deprecated. Use the JAX FFI instead.jax.lib.xla_client.dtype_to_etype,
jax.lib.xla_client.ops,
jax.lib.xla_client.shape_from_pyval, jax.lib.xla_client.PrimitiveType,
jax.lib.xla_client.Shape, jax.lib.xla_client.XlaBuilder, and
jax.lib.xla_client.XlaComputation have been deprecated. Use StableHLO
instead.Removals
jax.experimental.pallas.tpu.CostEstimate and
{func}jax.experimental.tpu.run_scoped. Both are now available in
{mod}jax.experimental.pallas.New functionality
pl.estimate_cost for automatically
constructing a kernel cost estimate from a JAX reference function.jax.experimental.host_callback has been deprecated since March 2024, with JAX version 0.4.26. Now we set the default value of the --jax_host_callback_…
New Functionality
jax.errors.JaxRuntimeError has been added as a public alias for the
formerly private XlaRuntimeError type.Breaking changes
jax_pmap_no_rank_reduction flag is set to True by default.
jax.experimental.host_callback has been deprecated since March 2024, with
JAX version 0.4.26. Now we set the default value of the
--jax_host_callback_legacy configuration value to True, which means that
if your code uses jax.experimental.host_callback APIs, those API calls
will be implemented in terms of the new jax.experimental.io_callback API.
If this breaks your code, for a very limited time, you can set the
--jax_host_callback_legacy to True. Soon we will remove that
configuration option, so you should instead transition to using the
new JAX callback APIs. See {jax-issue}#20385 for a discussion.Deprecations
jax.numpy.trim_zeros, non-arraylike arguments or arraylike
arguments with ndim != 1 are now deprecated, and in the future will result
in an error.jax.core.pp_* have been removed, after
being deprecated in JAX v0.4.30.jax.lib.xla_client.Device is deprecated; use jax.Device instead.jax.lib.xla_client.XlaRuntimeError has been deprecated. Use
jax.errors.JaxRuntimeError instead.Deletion:
jax.xla_computation is deleted. It's been 3 months since it's deprecation
in 0.4.30 JAX release.
Please use the AOT APIs to get the same functionality as jax.xla_computation.
jax.xla_computation(fn)(*args, **kwargs) can be replaced with
jax.jit(fn).lower(*args, **kwargs).compiler_ir('hlo')..out_info property of jax.stages.Lowered to get the
output information (like tree structure, shape and dtype).jax.xla_computation(fn, backend='tpu')(*args, **kwargs) with
jax.jit(fn).trace(*args, **kwargs).lower(lowering_platforms=('tpu',)).compiler_ir('hlo').jax.ShapeDtypeStruct no longer accepts the named_shape argument.
The argument was only used by xmap which was removed in 0.4.31.jax.tree.map(f, None, non-None), which previously emitted a
DeprecationWarning, now raises an error in a future version of jax. None
is only a tree-prefix of itself. To preserve the current behavior, you can
ask jax.tree.map to treat None as a leaf value by writing:
jax.tree.map(lambda x, y: None if x is None else f(x, y), a, b, is_leaf=lambda x: x is None).jax.sharding.XLACompatibleSharding has been removed. Please use
jax.sharding.Sharding.Bug fixes
jax.numpy.cumsum would produce incorrect outputs
if a non-boolean input was provided and dtype=bool was specified.jax.numpy.ldexp to get correct gradient.Changes
jax.experimental.pallas.debug_print no longer requires all arguments
to be scalars. The restrictions on the arguments are backend-specific:
Non-scalar arguments are currently only supported on GPU, when using Triton.jax.experimental.pallas.BlockSpec no longer supports the previously
deprecated argument order, where index_map comes before block_shape.Deprecations
jax.experimental.pallas.gpu submodule is deprecated to avoid
ambiguite with {mod}jax.experimental.pallas.mosaic_gpu. To use the
Triton backend import {mod}jax.experimental.pallas.triton.New functionality
jax.experimental.pallas.pallas_call now accepts scratch_shapes,
a PyTree specifying backend-specific temporary objects needed by the
kernel, for example, buffers, synchronization primitives etc.checkify.check can now be used to insert runtime asserts when
pallas_call is called with the pltpu.enable_runtime_assert(True) context
manager.This is a patch release on top of jax 0.4.32, that fixes two bugs found in that release.
This is a patch release on top of jax 0.4.32, that fixes two bugs found in that release.
A TPU-only data corruption bug was found in the version of libtpu pinned by
JAX 0.4.32, which manifested only if multiple TPU slices were present in the
same job, for example, if training on multiple v5e slices.
This release fixes that issue by pinning a fixed version of libtpu.
This release fixes an inaccurate result for F64 tanh on CPU (#23590).
WARNING: This release has been yanked from PyPI because of a data corruption bug on TPU if there are multiple TPU slices in the job
WARNING: This release has been yanked from PyPI because of a data corruption bug on TPU if there are multiple TPU slices in the job
Note: This release was yanked from PyPi because of a data corruption bug on TPU. See the 0.4.33 release notes for more details.
New Functionality
jax.extend.ffi.ffi_call and {func}jax.extend.ffi.ffi_lowering
to support the use of the new {ref}ffi-tutorial to interface with custom
C++ and CUDA code from JAX.Changes
jax_enable_memories flag is set to True by default.jax.numpy now supports v2023.12 of the Python Array API Standard.
See {ref}python-array-api for more information.jax.config.update('jax_cpu_enable_async_dispatch', False).jax.process_indices function to replace the
jax.process_indexs() function that was deprecated in JAX v0.2.13.numpy.fabs, jax.numpy.fabs has been
modified to no longer support complex dtypes.jax.tree_util.register_dataclass now checks that data_fields
and meta_fields includes all dataclass fields with init=True
and only them, if nodetype is a dataclass.jax.numpy functions now have full {class}~jax.numpy.ufunc
interfaces, including {obj}~jax.numpy.add, {obj}~jax.numpy.multiply,
{obj}~jax.numpy.bitwise_and, {obj}~jax.numpy.bitwise_or,
{obj}~jax.numpy.bitwise_xor, {obj}~jax.numpy.logical_and,
{obj}~jax.numpy.logical_and, and {obj}~jax.numpy.logical_and.
In future releases we plan to expand these to other ufuncs.jax.lax.optimization_barrier, which allows users to prevent
compiler optimizations such as common-subexpression elimination and to
control scheduling.Breaking changes
jax.extend.mlir.mhlo) has been removed. Use the
stablehlo dialect instead.Deprecations
jax.numpy.clip and {func}jax.numpy.hypot are
no longer allowed, after being deprecated since JAX v0.4.27.jax.lib.xla_bridge.xla_client: use {mod}jax.lib.xla_client directly.jax.lib.xla_bridge.get_backend: use {func}jax.extend.backend.get_backend.jax.lib.xla_bridge.default_backend: use {func}jax.extend.backend.default_backend.jax.experimental.array_api module is deprecated, and importing it is no
longer required to use the Array API. jax.numpy supports the array API
directly; see {ref}python-array-api for more information.jax.core.check_eqn, jax.core.check_type, and
jax.core.check_valid_jaxtype are now deprecated, and will be removed in
the future.jax.numpy.round_ has been deprecated, following removal of the corresponding
API in NumPy 2.0. Use {func}jax.numpy.round instead.jax.dlpack.from_dlpack is deprecated.
The argument to {func}jax.dlpack.from_dlpack should be an array from
another framework that implements the __dlpack__ protocol.Note: This release was yanked from PyPi because of a data corruption bug on TPU. See the 0.4.33 release notes for more details.
New Functionality
Added jax.extend.ffi.ffi_call() and jax.extend.ffi.ffi_lowering() to support the use of the new Foreign function interface (FFI) to interface with custom C++ and CUDA code from JAX.
Changes
jax_enable_memories flag is set to True by default.
jax.numpy now supports v2023.12 of the Python Array API Standard. See Python Array API standard for more information.
Computations on the CPU backend may now be dispatched asynchronously in more cases. Previously non-parallel computations were always dispatched synchronously. You can recover the old behavior by setting jax.config.update('jax_cpu_enable_async_dispatch', False) .
Added new jax.process_indices() function to replace the jax.process_indexs() function that was deprecated in JAX v0.2.13.
To align with the behavior of numpy.fabs , jax.numpy.fabs has been modified to no longer support complex dtypes .
jax.tree_util.register_dataclass now checks that data_fields and meta_fields includes all dataclass fields with init=True and only them, if nodetype is a dataclass.
Several jax.numpy functions now have full ufunc interfaces, including add , multiply , bitwise_and , bitwise_or , bitwise_xor , logical_and , logical_and , and logical_and . In future releases we plan to expand these to other ufuncs.
Added jax.lax.optimization_barrier() , which allows users to prevent compiler optimizations such as common-subexpression elimination and to control scheduling.
Breaking changes
The MHLO MLIR dialect ( jax.extend.mlir.mhlo ) has been removed. Use the stablehlo dialect instead.
Deprecations
Complex inputs to jax.numpy.clip() and jax.numpy.hypot() are no longer allowed, after being deprecated since JAX v0.4.27.
Deprecated the following APIs:
jax.lib.xla_bridge.xla_client : use jax.lib.xla_client directly.
jax.lib.xla_bridge.get_backend : use jax.extend.backend.get_backend() .
jax.lib.xla_bridge.default_backend : use jax.extend.backend.default_backend() .
The jax.experimental.array_api module is deprecated, and importing it is no longer required to use the Array API. jax.numpy supports the array API directly; see Python Array API standard for more information.
The internal utilities jax.core.check_eqn , jax.core.check_type , and jax.core.check_valid_jaxtype are now deprecated, and will be removed in the future.
jax.numpy.round_ has been deprecated, following removal of the corresponding API in NumPy 2.0. Use jax.numpy.round() instead.
Passing a DLPack capsule to jax.dlpack.from_dlpack() is deprecated. The argument to jax.dlpack.from_dlpack() should be an array from another framework that implements the dlpack protocol.
Changes
#22746).New functionality
{class}jax.experimental.pallas.BlockSpec now expects block_shape to be passed *before* index_map. The old argument order is deprecated and will be rem…
Deletion
shard_map as the replacement.Changes
jax.numpy.ceil, {func}jax.numpy.floor and {func}jax.numpy.trunc now return the output
of the same dtype as the input, i.e. no longer upcast integer or boolean inputs to floating point.libdevice.10.bc is no longer bundled with CUDA wheels. It must be
installed either as a part of local CUDA installation, or via NVIDIA's CUDA
pip wheels.jax.experimental.pallas.BlockSpec now expects block_shape to
be passed before index_map. The old argument order is deprecated and
will be removed in a future release.cuda(id=0) will now be CudaDevice(id=0).device property and to_device method to {class}jax.Array, as
part of JAX's Array API support.Deprecations
jax.core: removed canonicalize_shape,
dimension_as_value, definitely_equal, and symbolic_equal_dim.jax.experimental.jax2tf.convert with native_serialization=False
or enable_xla=False is now deprecated and this support will be removed in
a future version.
Native serialization has been the default since JAX 0.4.16 (September 2023).jax.random.shuffle has been removed;
instead use jax.random.permutation with independent=True.Changes
jax.experimental.pallas.BlockSpec now expects block_shape to
be passed before index_map. The old argument order is deprecated and
will be removed in a future release.jax.experimental.pallas.GridSpec does not have anymore the in_specs_tree,
and the out_specs_tree fields, and the in_specs and out_specs tree now
store the values as pytrees of BlockSpec. Previously, in_specs and
out_specs were flattened ({jax-issue}#22552).compute_index of {class}jax.experimental.pallas.GridSpec has
been removed because it is private. Similarly, the get_grid_mapping and
unzip_dynamic_bounds have been removed from BlockSpec ({jax-issue}#22593).#22275).
Padding in interpret mode will be with NaN, to help debug out-of-bounds
errors, but this behavior is not present when running in custom kernel mode,
and should not be depended on.jax.experimental.pallas.pallas. This is not possible anymore.New Functionality
pallas_grids_and_blockspecs.jax.experimental.pallas.pallas_call
API.lax.shift_right_arithmetic ({jax-issue}#22279)
and lax.erf_inv ({jax-issue}#22310).#22084).#22480)Added an API for exporting and serializing JAX functions. This used to exist in jax.experimental.export (which is being deprecated), and will now live…
Changes
jax.experimental.mesh_utils can now create an efficient mesh for TPU v5e.pip install jax, no extras required.jax.experimental.export (which is being deprecated),
and will now live in jax.export.
See the documentation.Deprecations
jax.core.pp_* are deprecated, and will be removed
in a future release.TypeError in a future JAX
release. This previously was the case, but there was an inadvertent regression in
the last several JAX releases.jax.experimental.export is deprecated. Use {mod}jax.export instead.
See the migration guide.x and y, x.astype(y) will raise a warning. To silence it use x.astype(y.dtype).jax.xla_computation is deprecated and will be removed in a future release.
Please use the AOT APIs to get the same functionality as jax.xla_computation.
jax.xla_computation(fn)(*args, **kwargs) can be replaced with
jax.jit(fn).lower(*args, **kwargs).compiler_ir('hlo')..out_info property of jax.stages.Lowered to get the
output information (like tree structure, shape and dtype).jax.xla_computation(fn, backend='tpu')(*args, **kwargs) with
jax.jit(fn).trace(*args, **kwargs).lower(lowering_platforms=('tpu',)).compiler_ir('hlo').jax.experimental.pallas.pallas_call in
interpret mode ({jax-issue}#21862).#21773).jax.tree.map(f, None, non-None) now emits a DeprecationWarning, and will raise an error in a future version of jax. None is only a tree-prefix of itse…
Bug fixes
Deprecations
jax.tree.map(f, None, non-None) now emits a DeprecationWarning, and will
raise an error in a future version of jax. None is only a tree-prefix of
itself. To preserve the current behavior, you can ask jax.tree.map to
treat None as a leaf value by writing:
jax.tree.map(lambda x, y: None if x is None else f(x, y), a, b, is_leaf=lambda x: x is None).Changes
pip install jax[cuda12]).jax.experimental.export API. It is not possible anymore to use
from jax.experimental.export import export, and instead you should use
from jax.experimental import export.
The removed functionality has been deprecated since 0.4.24.is_leaf argument to {func}jax.tree.all & {func}jax.tree_util.tree_all.Deprecations
jax.sharding.XLACompatibleSharding is deprecated. Please use
jax.sharding.Sharding.jax.experimental.Exported.in_shardings has been renamed as
jax.experimental.Exported.in_shardings_hlo. Same for out_shardings.
The old names will be removed after 3 months.jax.core: non_negative_dim, DimSize, Shapejax.lax: tie_injax.nn: normalizejax.interpreters.xla: backend_specific_translations,
translations, register_translation, xla_destructure,
TranslationRule, TranslationContext, XlaOp.tol argument of {func}jax.numpy.linalg.matrix_rank is being
deprecated and will soon be removed. Use rtol instead.rcond argument of {func}jax.numpy.linalg.pinv is being
deprecated and will soon be removed. Use rtol instead.jax.config submodule has been removed. To configure JAX
use import jax and then reference the config object via jax.config.jax.random APIs no longer accept batched keys, where previously
some did unintentionally. Going forward, we recommend explicit use of
{func}jax.vmap in such cases.jax.scipy.special.beta, the x and y parameters have been
renamed to a and b for consistency with other beta APIs.New Functionality
jax.experimental.Exported.in_shardings_jax to construct
shardings that can be used with the JAX APIs from the HloShardings
that are stored in the Exported objects.Fixes a memory corruption bug in the type name of Array and JIT Python objects in Python 3.10 or earlier.
Bug fixes
'+ptx84' is not a recognized feature for this target
under CUDA 12.4.Changes
Bug fixes
make_jaxpr that was breaking Equinox (#21116).Deprecations & removals
kind argument to {func}jax.numpy.sort and {func}jax.numpy.argsort
is now removed. Use stable=True or stable=False instead.get_compute_capability from the jax.experimental.pallas.gpu
module. Use the compute_capability attribute of a GPU device, returned
by {func}jax.devices or {func}jax.local_devices, instead.newshape argument to {func}jax.numpy.reshapeis being deprecated
and will soon be removed. Use shape instead.Changes
{func}jax.numpy.clip has a new argument signature: a, a_min, and a_max are deprecated in favor of x (positional only), min, and max ({jax-issue}20550)…
New Functionality
jax.numpy.unstack and {func}jax.numpy.cumulative_sum,
following their addition in the array API 2023 standard, soon to be
adopted by NumPy.jax_cpu_collectives_implementation to select the
implementation of cross-process collective operations used by the CPU backend.
Choices available are 'none'(default), 'gloo' and 'mpi' (requires jaxlib 0.4.26).
If set to 'none', cross-process collective operations are disabled.Changes
jax.pure_callback, {func}jax.experimental.io_callback
and {func}jax.debug.callback now use {class}jax.Array instead
of {class}np.ndarray. You can recover the old behavior by transforming
the arguments via jax.tree.map(np.asarray, args) before passing them
to the callback.complex_arr.astype(bool) now follows the same semantics as NumPy, returning
False where complex_arr is equal to 0 + 0j, and True otherwise.core.Token now is a non-trivial class which wraps a jax.Array. It could
be created and threaded in and out of computations to build up dependency.
The singleton object core.token has been removed, users now should create
and use fresh core.Token objects instead.jax.config.update('jax_threefry_gpu_kernel_lowering', True). If the new
default causes issues, please file a bug. Otherwise, we intend to remove
this flag in a future release.Deprecations & Removals
JAX_TRITON_COMPILE_VIA_XLA environment variable no longer has any effect.jax.numpy.clip has a new argument signature: a, a_min, and
a_max are deprecated in favor of x (positional only), min, and
max ({jax-issue}20550).device() method of JAX arrays has been removed, after being deprecated
since JAX v0.4.21. Use arr.devices() instead.initial argument to {func}jax.nn.softmax and {func}jax.nn.log_softmax
is deprecated; empty inputs to softmax are now supported without setting this.jax.jit, passing invalid static_argnums or static_argnames
now leads to an error rather than a warning.jax.numpy.hypot function now issues a deprecation warning when
passing complex-valued inputs to it. This will raise an error when the
deprecation is completed.jax.numpy.nonzero, {func}jax.numpy.where, and
related functions now raise an error, following a similar change in NumPy.jax_cpu_enable_gloo_collectives is deprecated.
Use jax.config.update('jax_cpu_collectives_implementation', 'gloo') instead.jax.Array.device_buffer and jax.Array.device_buffers methods have
been removed after being deprecated in JAX v0.4.22. Instead use
{attr}jax.Array.addressable_shards and {meth}jax.Array.addressable_data.condition, x, and y parameters of jax.numpy.where are now
positional-only, following deprecation of the keywords in JAX v0.4.21.jax.lax.linalg now must be
specified by keyword. Previously, this raised a DeprecationWarning.jax.numpy APIs,
including {func}~jax.numpy.apply_along_axis,
{func}~jax.numpy.apply_over_axes, {func}~jax.numpy.inner,
{func}~jax.numpy.outer, {func}~jax.numpy.cross,
{func}~jax.numpy.kron, and {func}~jax.numpy.lexsort.Bug fixes
jax.numpy.astype will now always return a copy when copy=True.
Previously, no copy would be made when the output array would have the same
dtype as the input array. This may result in some increased memory usage.
The default value is set to copy=False to preserve backwards compatibility.{func}jax.tree_map is deprecated; use jax.tree.map instead, or for backward compatibility with older JAX versions, use {func}jax.tree_util.tree_map.
New Functionality
jax.numpy.trapezoid, following the addition of this function in
NumPy 2.0.Changes
jax.numpy.geomspace now chooses the logarithmic spiral
branch consistent with that of NumPy 2.0.lax.rng_bit_generator, and in turn the 'rbg'
and 'unsafe_rbg' PRNG implementations, under jax.vmap has
changed so that
mapping over keys results in random generation only from the first
key in the batch.jax.random.key for construction of PRNG key arrays
rather than jax.random.PRNGKey.Deprecations & Removals
jax.tree_map is deprecated; use jax.tree.map instead, or for backward
compatibility with older JAX versions, use {func}jax.tree_util.tree_map.jax.clear_backends is deprecated as it does not necessarily do what
its name suggests and can lead to unexpected consequences, e.g., it will not
destroy existing backends and release corresponding owned resources. Use
{func}jax.clear_caches if you only want to clean up compilation caches.
For backward compatibility or you really need to switch/reinitialize the
default backend, use {func}jax.extend.backend.clear_backends.jax.experimental.maps module and jax.experimental.maps.xmap are
deprecated. Use jax.experimental.shard_map or jax.vmap with the
spmd_axis_name argument for expressing SPMD device-parallel computations.jax.experimental.host_callback module is deprecated.
Use instead the new JAX external callbacks.
Added JAX_HOST_CALLBACK_LEGACY flag to assist in the transition to the
new callbacks. See {jax-issue}#20385 for a discussion.jax.numpy.array_equal and {func}jax.numpy.array_equiv
that cannot be converted to a JAX array now results in an exception.jax_parallel_functions_output_gda has been removed.
This flag was long deprecated and did nothing; its use was a no-op.jax.interpreters.ad.config and
jax.interpreters.ad.source_info_util have now been removed. Use jax.config
and jax.extend.source_info_util instead.Several deprecated APIs in {mod}jax.interpreters.xla that were removed in v0.4.24 have been re-added in v0.4.25, including backend_specific_translatio…
New Features
x[True] or x[False].jax.tree module, with a more convenient interface for referencing functions
in {mod}jax.tree_util.jax.tree.transpose (i.e. {func}jax.tree_util.tree_transpose) now accepts
inner_treedef=None, in which case the inner treedef will be automatically inferred.Changes
JAX_TRITON_COMPILE_VIA_XLA environment variable to "0".jax.interpreters.xla that were removed in v0.4.24
have been re-added in v0.4.25, including backend_specific_translations,
translations, register_translation, xla_destructure, TranslationRule,
TranslationContext, and XLAOp. These are still considered deprecated, and
will be removed again in the future when better replacements are available.
Refer to {jax-issue}#19816 for discussion.Deprecations & Removals
jax.numpy.linalg.solve now shows a deprecation warning for batched 1D
solves with b.ndim > 1. In the future these will be treated as batched 2D
solves.api-compatibility).
These include
jax.config.config object anddefine_*_state and DEFINE_* methods of {data}jax.config.jax.config submodule via import jax.config is deprecated.
To configure JAX use import jax and then reference the config object
via jax.config.the core.non_negative_dim API (introduced recently) was deprecated and core.max_dim and core.min_dim were introduced ({jax-issue}#18953) to express ma…
Changes
rule parameter of mlir.register_lowering then add your
primitive to jax._src.dispatch.prim_requires_devices_during_lowering set.
This is needed because custom_partitioning and JAX callbacks need physical
devices to create Shardings during lowering.
This is a temporary state until we can create Shardings without physical
devices.jax.numpy.argsort and {func}jax.numpy.sort now support the stable
and descending arguments.jax.experimental.jax2tf and {mod}jax.experimental.export):
#19227)#19235) we now
consider dimension variables from different scopes to be different, even
if they have the same name. Symbolic expressions from different scopes
cannot interact, e.g., in arithmetic operations.
Scopes are introduced by {func}jax.experimental.jax2tf.convert,
{func}jax.experimental.export.symbolic_shape, {func}jax.experimental.export.symbolic_args_specs.
The scope of a symbolic expression e can be read with e.scope and passed
into the above functions to direct them to construct symbolic expressions in
a given scope.
See https://github.com/jax-ml/jax/blob/main/jax/experimental/jax2tf/README.md#user-specified-symbolic-constraints.#19231; note that this may result in user-visible behavior
changes)#19235).core.non_negative_dim API (introduced recently)
was deprecated and core.max_dim and core.min_dim were introduced
({jax-issue}#18953) to express max and min for symbolic dimensions.
You can use core.max_dim(d, 0) instead of core.non_negative_dim(d).shape_poly.is_poly_dim is deprecated in favor of export.is_symbolic_dim
({jax-issue}#19282).export.args_specs is deprecated in favor of export.symbolic_args_specs ({jax-issue}#19283`).shape_poly.PolyShape and jax2tf.PolyShape are deprecated, use
strings for polymorphic shapes specifications ({jax-issue}#19284).jax.experimental.jax2tf and {mod}jax.experimental.export.
See description of version numbers.jax.experimental.export. Instead of
from jax.experimental.export import export you should use now
from jax.experimental import export. The old way of importing will
continue to work for a deprecation period of 3 months.jax.scipy.stats.sem.jax.numpy.unique with return_inverse = True returns inverse indices
reshaped to the dimension of the input, following a similar change to
{func}numpy.unique in NumPy 2.0.jax.numpy.sign now returns x / abs(x) for nonzero complex inputs. This is
consistent with the behavior of {func}numpy.sign in NumPy version 2.0.jax.scipy.special.logsumexp with return_sign=True now uses the NumPy 2.0
convention for the complex sign, x / abs(x). This is consistent with the behavior
of {func}scipy.special.logsumexp in SciPy v1.13.Deprecations & Removals
api-compatibility).
This includes:
jax.core: TracerArrayConversionError,
TracerIntegerConversionError, UnexpectedTracerError,
as_hashable_function, collections, dtypes, lu, map,
namedtuple, partial, pp, ref, safe_zip, safe_map,
source_info_util, total_ordering, traceback_util, tuple_delete,
tuple_insert, and zip.jax.lax: dtypes, itertools, naryop, naryop_dtype_rule,
standard_abstract_eval, standard_naryop, standard_primitive,
standard_unop, unop, and unop_dtype_rule.jax.linear_util submodule and all its contents.jax.prng submodule and all its contents.jax.random: PRNGKeyArray, KeyArray, default_prng_impl,
threefry_2x32, threefry2x32_key, threefry2x32_p, rbg_key, and
unsafe_rbg_key.jax.tree_util: register_keypaths, AttributeKeyPathEntry, and
GetItemKeyPathEntry.jax.interpreters.xla: backend_specific_translations, translations,
register_translation, xla_destructure, TranslationRule, TranslationContext,
axis_groups, ShapedArray, ConcreteArray, AxisEnv, backend_compile,
and XLAOp.jax.numpy: NINF, NZERO, PZERO, row_stack, issubsctype,
trapz, and in1d.jax.scipy.linalg: tril and triu.PRNGKeyArray.unsafe_raw_array has been
removed. Use {func}jax.random.key_data instead.bool(empty_array) now raises an error rather than returning False. This
previously raised a deprecation warning, and follows a similar change in NumPy.jax.random: passing batched keys directly to random number generation functions,
such as {func}~jax.random.bits, {func}~jax.random.gamma, and others, is deprecated
and will emit a FutureWarning. Use jax.vmap for explicit batching.jax.lax.tie_in is deprecated: it has been a no-op since JAX v0.2.0.Fixed a bug that caused verbose logging from the GPU compiler during compilation.
The device_buffer and device_buffers properties of JAX arrays are deprecated. Explicit buffers have been replaced by the more flexible array sharding…
device_buffer and device_buffers properties of JAX arrays are deprecated.
Explicit buffers have been replaced by the more flexible array sharding interface,
but the previous outputs can be recovered this way:
arr.device_buffer becomes arr.addressable_data(0)arr.device_buffers becomes [x.data for x in arr.addressable_shards]The previously-deprecated sym_pos argument has been removed from {func}jax.scipy.linalg.solve. Use assume_a='pos' instead.
New Features
jax.nn.squareplus.Changes
jax.distributed.initialize().jax.distributed.initialize() in cloud TPU environments.Deprecations
sym_pos argument has been removed from
{func}jax.scipy.linalg.solve. Use assume_a='pos' instead.None to {func}jax.array or {func}jax.asarray, either directly or
within a list or tuple, is deprecated and now raises a {obj}FutureWarning.
It currently is converted to NaN, and in the future will raise a {obj}TypeError.condition, x, and y parameters to jax.numpy.where by
keyword arguments has been deprecated, to match numpy.where.jax.numpy.array_equal and {func}jax.numpy.array_equiv
that cannot be converted to a JAX array is deprecated and now raises a
{obj}DeprecationWaning. Currently the functions return False, in the future this
will raise an exception.device() method of JAX arrays is deprecated. Depending on the context, it may
be replaced with one of the following:
jax.Array.devices returns the set of all devices used by the array.jax.Array.sharding gives the sharding configuration used by the array.Fixed some type confusion between E4M3 and E5M2 float8 types.
Added {obj}jax.typing.DTypeLike, which can be used to annotate objects that are convertible to JAX dtypes.
New Features
jax.typing.DTypeLike, which can be used to annotate objects that
are convertible to JAX dtypes.jax.numpy.fill_diagonal.Changes
Bug fixes
A number of internal utilities and inadvertent exports in {mod}jax.lax have been deprecated, and will be removed in a future release.
Changes
cuda12_pip installation, NCCL should be installed
automatically. Currently, NCCL 2.16 or newer is required.jax.Array.item now supports optional index arguments.Deprecations
jax.lax have
been deprecated, and will be removed in a future release.
jax.lax.dtypes: use jax.dtypes instead.jax.lax.itertools: use itertools instead.naryop, naryop_dtype_rule, standard_abstract_eval, standard_naryop,
standard_primitive, standard_unop, unop, and unop_dtype_rule are
internal utilities, now deprecated without replacement.Bug fixes
Your coding agent can read these notes before it upgrades. Set up the MCP server →