diff --git a/distrax/_src/utils/jittable.py b/distrax/_src/utils/jittable.py index b11ffc61..150b5246 100644 --- a/distrax/_src/utils/jittable.py +++ b/distrax/_src/utils/jittable.py @@ -56,8 +56,8 @@ def tree_unflatten(cls, aux_data, children): def _is_jax_data(x): """Check whether `x` is an instance of a JAX-compatible type.""" - # If it's a tracer, then it's already been converted by JAX. - if isinstance(x, jax.core.Tracer): + # Tracers and AOT input descriptors have already been converted by JAX. + if isinstance(x, (jax.core.Tracer, jax.stages.ArgInfo)): return True # `jax.vmap` replaces vmappable leaves with `object()` during serialization. diff --git a/distrax/_src/utils/jittable_test.py b/distrax/_src/utils/jittable_test.py index c1ac29ae..fd5ea71c 100644 --- a/distrax/_src/utils/jittable_test.py +++ b/distrax/_src/utils/jittable_test.py @@ -94,6 +94,17 @@ def add_one_to_params(obj): add_one_to_params(DummyJittable(jnp.zeros((5,)))) add_one_to_params(DummyJittable(jnp.ones((5,)))) + def test_donatable_to_aot_compiled_function(self): + def get_params(obj): + return obj.data['params'] + + obj = DummyJittable(jnp.ones((5,))) + compiled = ( + jax.jit(get_params, donate_argnums=0).trace(obj).lower().compile() + ) + + np.testing.assert_array_equal(compiled(obj), jnp.ones((5,))) + def test_modifying_object_data_does_not_leak_tracers(self): @jax.jit def add_one_to_params(obj):