Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 5 additions & 21 deletions numpyro/ops/provenance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
Expand Down Expand Up @@ -52,32 +51,17 @@ 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 = {}
for n, v in kwargs.items():
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)


Expand Down
10 changes: 10 additions & 0 deletions test/ops/test_provenance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
Loading