$ the-wire · showcase
JAX FIXES ZERO-PROPAGATION BUG IN REVERSE-MODE AUTODIFF
By RepoJournal · Filed · About Google · Composed from the cited sources · methodology
JAX's core autodiff engine now correctly propagates zero tangents through remat operations, fixing a downstream Flax pattern that was leaking Tracers into production code.
The fix lands in jax.vjp and jax.linearize with a new in_nzs parameter, a tuple-tree of per-input nonzero-tangent flags that lets dead code paths stay dead [1]. This prevents the tracer leaks that were surfacing in certain gradient computations, a subtle but critical correctness issue for anyone chaining JAX's autodiff with complex control flow. The change is fully backward compatible: the returned callable still accepts tangent inputs corresponding to zeros, so existing code keeps working. In parallel, the team removed process_call as dead code [2], a cleanup that shrinks the autodiff surface and reduces maintenance burden. A separate quality-of-life fix caps array repr output at 1MiB to prevent logging overhead when debugging with large tensors [3].
Action items
- → Pull latest JAX and test gradient paths with remat and complex control flow google/jax [plan]
- → Monitor for tracer leaks in existing Flax models after upgrade google/jax [monitor]
References