diff --git a/oryx/core/interpreters/harvest.py b/oryx/core/interpreters/harvest.py index 7b0eb5e..a47da6c 100644 --- a/oryx/core/interpreters/harvest.py +++ b/oryx/core/interpreters/harvest.py @@ -814,7 +814,9 @@ def _get_harvest_metadata(closed_jaxpr, settings, *args): in_avals = jax_util.safe_map( lambda a: jax.typeof(a), flat_args) - pe.trace_to_jaxpr_dynamic(flat_fun, in_avals) + pe.trace_to_jaxpr( + flat_fun, FlatTree.flatten_args(*in_avals), debug_info=flat_fun.debug_info + ) metadata = aux() out_tree() return metadata @@ -1130,7 +1132,11 @@ def _reap_checkpoint_rule(trace, *invals, jaxpr, policy, prevent_cse, @lu.cache def _oryx_pjit_jaxpr(flat_fun, in_avals): - jaxpr, out_avals, consts = pe.trace_to_jaxpr_dynamic(flat_fun, in_avals) + closed_jaxpr, out_avals_tree = pe.trace_to_jaxpr( + flat_fun, FlatTree.flatten_args(*in_avals), debug_info=flat_fun.debug_info + ) + jaxpr, consts = closed_jaxpr.jaxpr, closed_jaxpr.consts + out_avals = out_avals_tree.tree if any(isinstance(c, jax_core.Tracer) for c in consts): jaxpr = pe.convert_constvars_jaxpr(jaxpr) jaxpr = pe.close_jaxpr(jaxpr) diff --git a/oryx/core/interpreters/propagate.py b/oryx/core/interpreters/propagate.py index 683dc92..6f3efe2 100644 --- a/oryx/core/interpreters/propagate.py +++ b/oryx/core/interpreters/propagate.py @@ -238,9 +238,12 @@ def identity_reducer(env, eqn, state, new_state): @lu.cache def _to_jaxpr(flat_fun, in_avals): - new_jaxpr, _, consts = pe.trace_to_jaxpr_dynamic(flat_fun, in_avals) - new_jaxpr = jex.core.ClosedJaxpr(new_jaxpr, consts) - return new_jaxpr + from jax._src.tree_util import FlatTree + + closed_jaxpr, _ = pe.trace_to_jaxpr( + flat_fun, FlatTree.flatten_args(*in_avals), debug_info=flat_fun.debug_info + ) + return closed_jaxpr def propagate(cell_type: Type[Cell], diff --git a/oryx/core/ppl/effect_handler.py b/oryx/core/ppl/effect_handler.py index 6b91739..ce970e9 100644 --- a/oryx/core/ppl/effect_handler.py +++ b/oryx/core/ppl/effect_handler.py @@ -254,9 +254,12 @@ def default_call_interpreter_rule(primitive: jax_core.CallPrimitive, @lu.cache def _to_jaxpr(flat_fun, in_avals): - new_jaxpr, _, consts = pe.trace_to_jaxpr_dynamic(flat_fun, in_avals) - new_jaxpr = jex.core.ClosedJaxpr(new_jaxpr, consts) - return new_jaxpr + from jax._src.tree_util import FlatTree + + closed_jaxpr, _ = pe.trace_to_jaxpr( + flat_fun, FlatTree.flatten_args(*in_avals), debug_info=flat_fun.debug_info + ) + return closed_jaxpr def _pjit_effect_handler_rule(rules, state, invals, **params): diff --git a/oryx/experimental/matching/jax_rewrite.py b/oryx/experimental/matching/jax_rewrite.py index 80681ae..992ad7d 100644 --- a/oryx/experimental/matching/jax_rewrite.py +++ b/oryx/experimental/matching/jax_rewrite.py @@ -697,9 +697,12 @@ def __hash__(self): @lu.cache def _to_jaxpr(flat_fun, in_avals): - new_jaxpr, _, consts = pe.trace_to_jaxpr_dynamic(flat_fun, in_avals) - new_jaxpr = jex.core.ClosedJaxpr(new_jaxpr, consts) - return new_jaxpr + from jax._src.tree_util import FlatTree + + closed_jaxpr, _ = pe.trace_to_jaxpr( + flat_fun, FlatTree.flatten_args(*in_avals), debug_info=flat_fun.debug_info + ) + return closed_jaxpr class PjitPrimitive(CallPrimitive):