diff --git a/numpyro/ops/provenance.py b/numpyro/ops/provenance.py index 3d9f2a876..523fd375f 100644 --- a/numpyro/ops/provenance.py +++ b/numpyro/ops/provenance.py @@ -2,7 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 import jax -from jax.api_util import debug_info, flatten_fun, shaped_abstractify +from jax.api_util import debug_info, flatten_fun from jax.extend.core import Literal from jax.extend.core.primitives import call_p, closed_call_p, jit_p @@ -11,7 +11,6 @@ except ImportError: xla_pmap_p = None import jax.extend.linear_util as lu -from jax.interpreters.partial_eval import trace_to_jaxpr_dynamic # Adapted from definition in jax v0.5.0 @@ -52,24 +51,9 @@ def eval_provenance(fn, **kwargs): else {} ) wrapped_fun, out_tree = flatten_fun(lu.wrap_init(fn, **fn_info), in_tree) - # Abstract eval to get output pytree - avals = _safe_map(shaped_abstractify, args) - # Note: we split out the process of abstract evaluation and provenance tracking - # for simplicity. In principle, they can be merged so that we only need to walk - # through the equations once. - - wrapped_info = ( - dict( - debug_info=debug_info( - "eval_provenance wrapped", wrapped_fun.call_wrapped, args, {} - ) - ) - if debug_info is not None - else {} - ) - jaxpr, avals_out, _ = trace_to_jaxpr_dynamic( - lu.wrap_init(wrapped_fun.call_wrapped, {}, **wrapped_info), avals - ) + # Use JAX's public tracing API so closed-over constants remain separate from + # the explicit inputs whose provenance is tracked. + jaxpr = jax.make_jaxpr(wrapped_fun.call_wrapped)(*args) # get provenances of flatten kwargs aval_kwargs = {} @@ -77,7 +61,7 @@ def eval_provenance(fn, **kwargs): aval_kwargs[n] = jax.tree.map(lambda _: frozenset({n}), v) provenance_inputs, _ = jax.tree.flatten(((), aval_kwargs)) - provenance_outputs = track_deps_jaxpr(jaxpr, provenance_inputs) + provenance_outputs = track_deps_jaxpr(jaxpr.jaxpr, provenance_inputs) return jax.tree.unflatten(out_tree(), provenance_outputs) diff --git a/test/ops/test_provenance.py b/test/ops/test_provenance.py index 4c88d973c..b8f62979a 100644 --- a/test/ops/test_provenance.py +++ b/test/ops/test_provenance.py @@ -20,9 +20,13 @@ from numpyro.ops.provenance import eval_provenance _JAX_SUBFUNS_API = Version(jax.__version__) >= Version("0.9.2") +_JAX_CALL_JAXPR_API = Version(jax.__version__) >= Version("0.11.1") def _call_bind(primitive, fn, *args): + if _JAX_CALL_JAXPR_API: + call_jaxpr = jax.make_jaxpr(fn.call_wrapped)(*args) + return primitive.bind(*args, call_jaxpr=call_jaxpr) if _JAX_SUBFUNS_API: return primitive.bind(*args, subfuns=(fn,)) return primitive.bind(fn, *args) @@ -57,6 +61,12 @@ def f(x): assert eval_provenance(f, x=3) == {"x"} +def test_provenance_closed_over_const(): + const = jnp.arange(5.0) + + assert eval_provenance(lambda x: const * x, x=1.0) == {"x"} + + def test_provenance_fori(): def f(x, y, z): del z