v0.4.34
JAX v0.4.34
Added 2
- Wheels for Python 3.13 are now available
- jax.errors.JaxRuntimeError has been added as a public alias for XlaRuntimeError
Changed 4
- jax_pmap_no_rank_reduction flag is now set to True by default
- array[0] on a pmap result now introduces a reshape; use array[0:1] instead
- Per-shard shape accessed via addressable_shards or addressable_data(0) now has a leading (1, ...) dimension
- Default value of --jax_host_callback_legacy configuration is now True, implementing jax.experimental.host_callback APIs in terms of jax.experimental.io_callback
Fixed 2
- Fixed a bug where jax.numpy.cumsum would produce incorrect outputs if a non-boolean input was provided and dtype=bool was specified
- Fixed implementation of jax.numpy.ldexp to get correct gradient
Removed 5
- Internal pretty-printing tools jax.core.pp_* have been removed
- jax.xla_computation has been deleted
- jax.ShapeDtypeStruct no longer accepts the named_shape argument
- jax.tree.map(f, None, non-None) now raises an error instead of emitting a DeprecationWarning
- jax.sharding.XLACompatibleSharding has been removed; use jax.sharding.Sharding instead
Deprecated 3
- Non-arraylike arguments or arraylike arguments with ndim != 1 in jax.numpy.trim_zeros are now deprecated
- jax.lib.xla_client.Device is deprecated; use jax.Device instead
- jax.lib.xla_client.XlaRuntimeError is deprecated; use jax.errors.JaxRuntimeError instead
-
New Functionality
- This release includes wheels for Python 3.13. Free-threading mode is not yet supported.
jax.errors.JaxRuntimeErrorhas been added as a public alias for the formerly privateXlaRuntimeErrortype.
-
Breaking changes
jax_pmap_no_rank_reductionflag is set toTrueby default.array[0]on a pmap result now introduces a reshape (usearray[0:1]instead).- The per-shard shape (accessable via
jax_array.addressable_shardsorjax_array.addressable_data(0))now has a leading(1, ...). Update code that directly accesses shards accordingly. The rank of the per-shard-shape now matches that of the global shape which is the same behavior as jit. This avoids costly reshapes when passing results from pmap into jit.
jax.experimental.host_callbackhas been deprecated since March 2024, with JAX version 0.4.26. Now we set the default value of the--jax_host_callback_legacyconfiguration value toTrue, which means that if your code usesjax.experimental.host_callbackAPIs, those API calls will be implemented in terms of the newjax.experimental.io_callbackAPI. If this breaks your code, for a very limited time, you can set the--jax_host_callback_legacytoTrue. Soon we will remove that configuration option, so you should instead transition to using the new JAX callback APIs. See #20385 for a discussion.
-
Deprecations
- In
jax.numpy.trim_zeros, non-arraylike arguments or arraylike arguments withndim != 1are now deprecated, and in the future will result in an error. - Internal pretty-printing tools
jax.core.pp_*have been removed, after being deprecated in JAX v0.4.30. jax.lib.xla_client.Deviceis deprecated; usejax.Deviceinstead.jax.lib.xla_client.XlaRuntimeErrorhas been deprecated. Usejax.errors.JaxRuntimeErrorinstead.
- In
-
Deletion:
jax.xla_computationis deleted. It has been 3 months since its deprecation in 0.4.30 JAX release. Please use the AOT APIs to get the same functionality asjax.xla_computation.jax.xla_computation(fn)(*args, **kwargs)can be replaced withjax.jit(fn).lower(*args, **kwargs).compiler_ir('hlo').- You can also use
.out_infoproperty ofjax.stages.Loweredto get the output information (like tree structure, shape and dtype). - For cross-backend lowering, you can replace
jax.xla_computation(fn, backend='tpu')(*args, **kwargs)withjax.jit(fn).trace(*args, **kwargs).lower(lowering_platforms=('tpu',)).compiler_ir('hlo').
jax.ShapeDtypeStructno longer accepts thenamed_shapeargument. The argument was only used byxmapwhich was removed in 0.4.31.jax.tree.map(f, None, non-None), which previously emitted aDeprecationWarning, now raises an error.Noneis only a tree-prefix of itself. To preserve the current behavior, you can askjax.tree.mapto treatNoneas 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.XLACompatibleShardinghas been removed. Please usejax.sharding.Sharding.
-
Bug fixes
- Fixed a bug where
jax.numpy.cumsumwould produce incorrect outputs if a non-boolean input was provided anddtype=boolwas specified. - Edit implementation of
jax.numpy.ldexpto get correct gradient.
- Fixed a bug where