v0.4.13
JAX release v0.4.13
Added 2
- Added a warning if a non-allowlisted jaxlib plugin is in use
- Added jax.tree_util.tree_leaves_with_path
Changed 3
- jax.jit now allows None to be passed to in_shardings and out_shardings, with in_shardings marked as replicated and out_shardings determined by the XLA GSPMD partitioner
- jax.experimental.pjit.pjit now allows None to be passed to in_shardings and out_shardings with mesh-dependent semantics
- Executable.cost_analysis() now works on Cloud TPU
Fixed 1
- Fixed incorrect wheel name in CUDA 12 releases; the correct wheel is named cudnn89 instead of cudnn88
Deprecated 1
- The native_serialization_strict_checks parameter to jax.experimental.jax2tf.convert is deprecated in favor of native_serializaation_disabled_checks
NOTE: This is the last JAX release that will include Python 3.8 support
-
Changes
jax.jitnow allowsNoneto be passed toin_shardingsandout_shardings. The semantics are as follows:- For in_shardings, JAX will mark is as replicated but this behavior can change in the future.
- For out_shardings, we will rely on the XLA GSPMD partitioner to determine the output shardings.
jax.experimental.pjit.pjitalso allowsNoneto be passed toin_shardingsandout_shardings. The semantics are as follows:- If the mesh context manager is not provided, JAX has the freedom to
choose whatever sharding it wants.
- For in_shardings, JAX will mark is as replicated but this behavior can change in the future.
- For out_shardings, we will rely on the XLA GSPMD partitioner to determine the output shardings.
- If the mesh context manager is provided, None will imply that the value will be replicated on all devices of the mesh.
- If the mesh context manager is not provided, JAX has the freedom to
choose whatever sharding it wants.
- Executable.cost_analysis() works on Cloud TPU
- Added a warning if a non-allowlisted
jaxlibplugin is in use. - Added
jax.tree_util.tree_leaves_with_path.
-
Bug fixes
- Fixed incorrect wheel name in CUDA 12 releases (#16362); the correct wheel
is named
cudnn89instead ofcudnn88.
- Fixed incorrect wheel name in CUDA 12 releases (#16362); the correct wheel
is named
-
Deprecations
- The
native_serialization_strict_checksparameter to {func}jax.experimental.jax2tf.convertis deprecated in favor of the newnative_serializaation_disabled_checks({jax-issue}#16347).
- The