$ the-wire · showcase
JAX pins down DCE jaxpr identity, fixes i0/i1 gradients at zero
By RepoJournal · Filed · About Google · Composed from the cited sources · methodology
JAX fixed a nondeterminism bug where lowering the same jitted function could emit different StableHLO depending on what the process had traced earlier, and corrected higher-order gradients of i0 and i1 at 0.0.
Lowering the same jitted function used to emit different StableHLO depending on trace history. The MLIR lowering cache dedups repeated subcomputations by jaxpr object identity, but DCE always rebuilt a new jaxpr even when it eliminated nothing, so that identity was only stable while the right weakref LRU cache entries stayed warm. The fix makes no-op DCE return the input jaxpr [1], and the PR notes it fixes the reported repro while one narrower subclass remains open.
The JVP rule for `jax.numpy.i0` used to return a constant zero at `x=0.0`, producing incorrect higher-order derivatives; it now uses a maclaurin series approach for the JVP of i1, similar to the approach used for `sinc` [2]. JAX also added a `lax.one_minus_square` primitive to preserve accuracy of `1 - x^2` when evaluating near `+/-1` and differentiating near `0`, addressing the tension in expressions such as the derivatives of `tanh`, `atanh`, `asin`, and `acos` [3]. Separately, Ref buffers now keep their memory space constraints during state discharge [4].
In python-genai, the automatic function calling loop no longer runs functions once its budget is spent [6], closing a path where calls continued past the configured limit. RetrievalCallStep and RetrievalResultStep were added to the interactions schema and SDKs [5].
google-cloud-python landed Tier 3 client method span wrapping in `google.api_core.gapic_v1.method._GapicCallable.__call__`, starting an OpenTelemetry `SpanKind.CLIENT` span around the high-level GAPIC SDK method call and setting `rpc.system = "grpc"`, `rpc.service`, and `rpc.method`, encompassing client preparation, retry loops, timeouts, and error handling [7]. `google-cloud-spanner` moves to its own release PR for independent release management [8].
Action items
- → Re-run numerical tests on any code differentiating jax.numpy.i0 or i1 at 0.0 google/jax [plan]
- → Audit automatic function calling budgets in python-genai before upgrading googleapis/python-genai [monitor]
- → Plan for GAPIC method spans appearing in OpenTelemetry traces after upgrading google-api-core googleapis/google-cloud-python [monitor]
- → Update Spanner release pinning to account for its move to individual releases googleapis/google-cloud-python [plan]
References
- [1] Make no-op DCE return the input jaxpr, so lowering does not depend on trace history ↗ google/jax
- [2] [autodiff] fix first, second, and higher-order gradients of i0 and i1 at 0.0 ↗ google/jax
- [3] Add `lax.one_minus_square` primitive to preserve accuracy of `1 - x^2` when evaluating near `+/-1` and differentiating near `0`. ↗ google/jax
- [4] Preserve memory space constraints on Ref buffers during state discharge. ↗ google/jax
- [5] feat: Add RetrievalCallStep and RetrievalResultStep to interactions schema and SDKs. ↗ googleapis/python-genai
- [6] fix: do not run functions once the automatic function calling budget is spent ↗ googleapis/python-genai
- [7] feat(gapic): add OpenTelemetry T3 client method span wrapping in gapic_v1.method (D) ↗ googleapis/google-cloud-python
- [8] chore(spanner): move to individual releases ↗ googleapis/google-cloud-python