diff --git a/flax/nnx/transforms/compilation.py b/flax/nnx/transforms/compilation.py index 1b55009a6..c047765d0 100644 --- a/flax/nnx/transforms/compilation.py +++ b/flax/nnx/transforms/compilation.py @@ -20,6 +20,7 @@ import operator import typing as tp +import numpy as np import jax from jax.sharding import AbstractMesh, Mesh, PartitionSpec @@ -354,6 +355,17 @@ def jit( support shared ``Variable`` references or returning mutable array references from the jitted function. + .. note:: + **Captured Variables (hijax only).** When ``flax_hijax_variable`` + is enabled, ``fun`` may close over NNX ``Variable`` objects. + ``nnx.jit`` discovers them via jaxpr tracing (they appear as + ``AbstractVariable`` constants), promotes them to explicit inputs + for buffer donation, and relies on hijax ``MutableHiType`` for + mutation write-back. Without ``flax_hijax_variable``, closure + Variables are treated as opaque constants by JAX — mutations will + **not** propagate back. Pass Variables as explicit arguments + instead. + Returns: A wrapped version of ``fun``, set up for just-in-time compilation. """ @@ -527,6 +539,8 @@ def __init__( graph: bool, update_shardings: tuple[tp.Any, ...], ): + from flax.configurations import config as flax_config + functools.update_wrapper(self, fun) self.fun: tp.Callable[P, R] = fun self.out_shardings = out_shardings @@ -542,36 +556,151 @@ def __init__( else: self.in_shardings = in_shardings - donate_argnums_set = frozenset( - (donate_argnums,) if isinstance(donate_argnums, int) - else donate_argnums or () - ) - donate_argnames_set = frozenset( - (donate_argnames,) if isinstance(donate_argnames, str) - else donate_argnames or () + # ponytail: captures only supported via hijax jaxpr path + if not graph and flax_config.flax_hijax_variable: + self._use_jaxpr = True + self._jaxpr_cache: dict[tp.Any, tuple[tp.Any, list[tp.Any], int]] = {} + self._static_argnums = ( + (static_argnums,) if isinstance(static_argnums, int) + else tuple(static_argnums) if static_argnums + else () + ) + self._static_argnames = ( + (static_argnames,) if isinstance(static_argnames, str) + else tuple(static_argnames) if static_argnames + else () + ) + self._donate_argnums = donate_argnums + self._donate_argnames = donate_argnames + self._keep_unused = keep_unused + self._device = device + self._backend = backend + self._inline = inline + self._in_shardings_raw = in_shardings + else: + self._use_jaxpr = False + self._jaxpr_cache = {} + self._static_argnums = () + self._static_argnames = () + + donate_argnums_set = frozenset( + (donate_argnums,) if isinstance(donate_argnums, int) + else donate_argnums or () + ) + donate_argnames_set = frozenset( + (donate_argnames,) if isinstance(donate_argnames, str) + else donate_argnames or () + ) + + self._jitted_fn_legacy = jax.jit( + SimpleJitFn( + fun, + self.in_shardings, + out_shardings, + donate_argnums_set, + donate_argnames_set, + graph, + tuple(update_shardings), + ), + in_shardings=in_shardings, + out_shardings=(out_shardings, update_shardings), + static_argnums=static_argnums, + static_argnames=static_argnames, + donate_argnums=donate_argnums, + donate_argnames=donate_argnames, + keep_unused=keep_unused, + device=device, + backend=backend, + inline=inline, + ) + + @property + def jitted_fn(self): + """Access the underlying jitted function (for cache inspection etc).""" + if self._jaxpr_cache: + return next(iter(self._jaxpr_cache.values()))[0] + if hasattr(self, '_jitted_fn_legacy'): + return self._jitted_fn_legacy + raise AttributeError('jitted_fn not available before first call') + + def _get_jaxpr_jitted(self, args, kwargs): + """Lazily trace, transform, and cache a jitted eval function. + + On the first call (per unique static-arg values), traces the user + function with ``jax.make_jaxpr``, identifies ``AbstractVariable`` + constants (captured closure Variables), promotes them to explicit + jaxpr inputs, and wraps the result with ``jax.jit``. Subsequent + calls with the same static values reuse the cached version. + """ + from flax.nnx.variablelib import AbstractVariable + from jax._src import core as jax_core + + # Cache key: static arg values (different statics → different jaxprs) + static_key: tuple[tp.Any, ...] = () + for i in self._static_argnums: + if i < len(args): + static_key += (i, args[i]) + for name in self._static_argnames: + if name in kwargs: + static_key += (name, kwargs[name]) + + if static_key in self._jaxpr_cache: + return self._jaxpr_cache[static_key] + + # 1. Trace the function → ClosedJaxpr + all_args = (*self.partial_args, *args) + closed, out_shapes = jax.make_jaxpr( + self.fun, + static_argnums=tuple( + i + len(self.partial_args) for i in self._static_argnums + ) or None, + return_shape=True, + )(*all_args, **kwargs) + _, out_treedef = jax.tree_util.tree_flatten(out_shapes) + + # 2. Identify AbstractVariable constants → promote to inputs + jaxpr = closed.jaxpr + cvars = np.array(jaxpr.constvars, dtype=object) + consts = np.array(closed.consts, dtype=object) + mask = np.array([isinstance(cv.aval, AbstractVariable) for cv in cvars]) + + new_jaxpr = jaxpr.replace( + constvars=list(cvars[~mask]), + invars=list(cvars[mask]) + list(jaxpr.invars), ) - self.jitted_fn = jax.jit( - SimpleJitFn( - fun, - self.in_shardings, - out_shardings, - donate_argnums_set, - donate_argnames_set, - graph, - tuple(update_shardings), - ), - in_shardings=in_shardings, - out_shardings=(out_shardings, update_shardings), - static_argnums=static_argnums, - static_argnames=static_argnames, - donate_argnums=donate_argnums, - donate_argnames=donate_argnames, - keep_unused=keep_unused, - device=device, - backend=backend, - inline=inline, + + # 3. Build eval function (no closure Variables) + n_var = int(mask.sum()) + + def eval_fn(*all_flat): + results = jax_core.eval_jaxpr(new_jaxpr, list(consts[~mask]), *all_flat) + return jax.tree.unflatten(out_treedef, results) + + # 4. Compute donate_argnums: Variable inputs (0..n_var-1) always + # donated, plus any user-specified donate_argnums (offset). + donate_list = list(range(n_var)) + if self._donate_argnums is not None: + user_donate = ( + (self._donate_argnums,) if isinstance(self._donate_argnums, int) + else self._donate_argnums + ) + for d in user_donate: + donate_list.append(n_var + d) + + jitted = jax.jit( + eval_fn, + donate_argnums=tuple(donate_list) or None, + donate_argnames=self._donate_argnames or None, + keep_unused=self._keep_unused, + device=self._device, + backend=self._backend, + inline=self._inline, ) + entry = (jitted, list(consts[mask]), n_var) + self._jaxpr_cache[static_key] = entry + return entry + def _maybe_to_tree(self, args, kwargs): if self.graph: if self.in_shardings is not None and isinstance(self.in_shardings, (tuple, list)): @@ -593,41 +722,100 @@ def _maybe_from_tree(self, out): out = extract.from_tree2(out) return out + def _check_no_capture_aliases(self, var_consts, args, kwargs): + """Check that no captured Variable also appears in user args.""" + captured_ids = {id(v) for v in var_consts} + is_leaf = lambda x: isinstance(x, variablelib.Variable) + for path, leaf in jax.tree_util.tree_leaves_with_path( + (args, kwargs), is_leaf=is_leaf + ): + if isinstance(leaf, variablelib.Variable) and id(leaf) in captured_ids: + path_str = jax.tree_util.keystr(path) + raise ValueError( + f'Duplicate {leaf}\nfound at paths:\n\n' + f' - captured_args/...\n' + f' - {path_str}\n\n' + f'nnx.jit with graph_updates=False does not support ' + 'Variable aliasing (duplicate inputs, duplicate outputs, or ' + 'input Variables returned as outputs). ' + f'Consider the following options:\n\n' + f'1. Remove the duplicate Variables.\n' + f'2. Create new Variables via nnx.clone() and use those instead.\n' + f'3. Enable graph mode and graph updates by passing graph=True and ' + f'graph_updates=True to jit\n\n' + f' nnx.jit(..., graph=True, graph_updates=True)\n\n' + f'4. Use nnx.compat.jit (sets graph and graph_updates to True ' + f'automatically)\n\n' + f' nnx.compat.jit(...)' + ) + def __call__(self, *args: P.args, **kwargs: P.kwargs) -> R: - args = (*self.partial_args, *args) # type: ignore[assignment] - args, kwargs = self._maybe_to_tree(args, kwargs) - variables = extract.check_no_aliases('jit', args=args, kwargs=kwargs) - out, updates = self.jitted_fn(*args, **kwargs) - extract.apply_updates(variables, updates) - return self._maybe_from_tree(out) + if self._use_jaxpr: + jitted, var_consts, n_var = self._get_jaxpr_jitted(args, kwargs) + self._check_no_capture_aliases(var_consts, args, kwargs) + da, dk = self._strip_statics(args, kwargs) + return jitted(*var_consts, *self.partial_args, *da, **dk) + else: + all_args: tuple[tp.Any, ...] = (*self.partial_args, *args) + all_args, kwargs = self._maybe_to_tree(all_args, kwargs) + variables = extract.check_no_aliases( + 'jit', args=all_args, kwargs=kwargs + ) + out, updates = self._jitted_fn_legacy(*all_args, **kwargs) + extract.apply_updates(variables, updates) + return self._maybe_from_tree(out) def __get__(self, obj, objtype=None): if obj is None: return self return functools.partial(self, obj) + def _strip_statics(self, args, kwargs): + dynamic_args = tuple( + a for i, a in enumerate(args) if i not in self._static_argnums + ) + dynamic_kwargs = { + k: v for k, v in kwargs.items() + if k not in self._static_argnames + } + return dynamic_args, dynamic_kwargs + def eval_shape(self, *args, **kwargs): + if self._use_jaxpr: + jitted, var_consts, _ = self._get_jaxpr_jitted(args, kwargs) + da, dk = self._strip_statics(args, kwargs) + return jitted.eval_shape(*var_consts, *self.partial_args, *da, **dk) args = (*self.partial_args, *args) args, kwargs = self._maybe_to_tree(args, kwargs) if not self.graph: extract.check_no_aliases('jit', args=args, kwargs=kwargs) - out, _ = self.jitted_fn.eval_shape(*args, **kwargs) + out, _ = self._jitted_fn_legacy.eval_shape(*args, **kwargs) return self._maybe_from_tree(out) def trace(self, *args, **kwargs): + if self._use_jaxpr: + jitted, var_consts, _ = self._get_jaxpr_jitted(args, kwargs) + da, dk = self._strip_statics(args, kwargs) + traced = jitted.trace(*var_consts, *self.partial_args, *da, **dk) + return SimpleTraced(traced, self) args = (*self.partial_args, *args) args, kwargs = self._maybe_to_tree(args, kwargs) if not self.graph: extract.check_no_aliases('jit', args=args, kwargs=kwargs) - traced = self.jitted_fn.trace(*args, **kwargs) + traced = self._jitted_fn_legacy.trace(*args, **kwargs) return SimpleTraced(traced, self) def lower(self, *args, **kwargs): + if self._use_jaxpr: + jitted, var_consts, _ = self._get_jaxpr_jitted(args, kwargs) + da, dk = self._strip_statics(args, kwargs) + lowered = jitted.lower(*var_consts, *self.partial_args, *da, **dk) + return SimpleLowered(lowered, self) args = (*self.partial_args, *args) args, kwargs = self._maybe_to_tree(args, kwargs) if not self.graph: extract.check_no_aliases('jit', args=args, kwargs=kwargs) - lowered = self.jitted_fn.lower(*args, **kwargs) + lowered = self._jitted_fn_legacy.lower(*args, **kwargs) return SimpleLowered(lowered, self) @@ -1335,7 +1523,8 @@ def call(*args, **kwargs): def __call__(self, *args, **kwargs): args = (*self.jit_wrapped.partial_args, *args) args, kwargs = self.jit_wrapped._maybe_to_tree(args, kwargs) - variables = extract.check_no_aliases('jit', args=args, kwargs=kwargs) + variables = extract.check_no_aliases( + 'jit', args=args, kwargs=kwargs) out, updates = self.compiled(*args, **kwargs) extract.apply_updates(variables, updates) return self.jit_wrapped._maybe_from_tree(out) diff --git a/tests/nnx/spmd_test.py b/tests/nnx/spmd_test.py index 5f075ce77..5c7fde199 100644 --- a/tests/nnx/spmd_test.py +++ b/tests/nnx/spmd_test.py @@ -605,6 +605,30 @@ def test_get_abstract_no_sharding_metadata(self): getattr(abs_model.kernel.get_value(), 'sharding', None) ) + def test_jit_closure_partition_spec_in_shardings(self): + """A PartitionSpec in_shardings applies to the arg but not captures.""" + mesh = jax.make_mesh((4,), ('data',)) + data_sharding = NamedSharding(mesh, P('data')) + + with jax.set_mesh(mesh): + model = nnx.Linear(2, 3, rngs=nnx.Rngs(0)) + observed_specs = [] + + def callback(sharding): + observed_specs.append(sharding.spec) + + @nnx.jit(in_shardings=(P('data'),), graph=False) + def forward(x): + jax.debug.inspect_array_sharding(model.kernel[...], callback=callback) + jax.debug.inspect_array_sharding(x, callback=callback) + return model(x) + + x = jax.device_put(jnp.ones((4, 2)), data_sharding) + forward(x) + self.assertEqual(observed_specs[0], P()) + self.assertEqual(observed_specs[1], P('data')) + + def has_sharding_spec(array): sharding = array.sharding if hasattr(sharding, 'spec'): diff --git a/tests/nnx/transforms_test.py b/tests/nnx/transforms_test.py index b1d75bc73..b2fb7a36b 100644 --- a/tests/nnx/transforms_test.py +++ b/tests/nnx/transforms_test.py @@ -16,7 +16,7 @@ os.environ['XLA_FLAGS'] = '--xla_force_host_platform_device_count=4' import dataclasses -from functools import partial +from functools import partial, wraps import typing as tp from absl.testing import absltest @@ -8017,5 +8017,198 @@ def forward_block(rng_state, rest_state, x): assert not jnp.allclose(y, y2) +class TestClosureCapture(parameterized.TestCase): + """Tests for auto-capture of closure Variables via hijax jaxpr path.""" + + def test_jit_closure_state_update(self): + """Mutations to captured Variables propagate back.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + count = nnx.Variable(jnp.array(0)) + + @nnx.jit(graph=False) + def forward(): + count[...] += 1 + + forward() + self.assertEqual(count[...], 1) + forward() + self.assertEqual(count[...], 2) + + def test_jit_closure_pytree_state_update(self): + """Mutations to captured pytrees of Variables propagate back.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + count = [nnx.Variable(jnp.array(0))] + + @nnx.jit(graph=False) + def forward(): + count[0][...] += 1 + + forward() + self.assertEqual(count[0][...], 1) + forward() + self.assertEqual(count[0][...], 2) + + def test_jit_closure_duplicate_raises(self): + """Passing a captured Variable as an explicit arg raises a duplicate error.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + count = nnx.Variable(jnp.array(0)) + + @nnx.jit(graph=False) + def forward(x): + count[...] += 1 + + with self.assertRaisesRegex(ValueError, 'Duplicate'): + forward(count) + + def test_jit_closure_nested(self): + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + count = nnx.Variable(jnp.array(0)) + + def forward(x): + count[...] += 1 + + f = partial(forward, 2) + nnx.jit(f, graph_updates=False)() + self.assertEqual(count[...], 1) + + def test_jit_closure_in_decorator(self): + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + count = nnx.Variable(jnp.array(0)) + + def my_decorator(f): + @wraps(f) + def wrapper(x): + count[...] += 1 + f(x) + return wrapper + + def silly(x): + count[...] += 1 + + nnx.jit(my_decorator(silly), graph_updates=False)(2) + self.assertEqual(count[...], 2) + + def test_jit_closure_no_recompilation(self): + """nnx.jit with a captured variable does not recompile on repeated calls.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + model = nnx.Linear(2, 3, rngs=nnx.Rngs(0)) + + @nnx.jit(graph=False) + def forward(x): + return model(x) + + x = jnp.ones((4, 2)) + forward(x) + cache_size_after_first = forward.jitted_fn._cache_size() + model.kernel[...] = jnp.zeros((2,3)) + forward(x) + cache_size_after_second = forward.jitted_fn._cache_size() + + self.assertEqual(cache_size_after_first, cache_size_after_second, + 'nnx.jit recompiled on the second call with a closure capture') + + def test_jit_closure_static_argnums(self): + """static_argnums indices refer to user args, not shifted by captures.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + model = nnx.Linear(2, 3, rngs=nnx.Rngs(0)) + + @nnx.jit(graph=False, static_argnums=1) + def forward(x, mode): + if mode == 'train': + return model(x) + return model(x) * 0 + + x = jnp.ones((4, 2)) + y_train = forward(x, 'train') + y_eval = forward(x, 'eval') + np.testing.assert_array_equal(y_eval, jnp.zeros_like(y_eval)) + self.assertFalse(jnp.allclose(y_train, jnp.zeros_like(y_train))) + + def test_jit_closure_static_argnums_sequence(self): + """static_argnums as a sequence is offset correctly with captures.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + model = nnx.Linear(2, 3, rngs=nnx.Rngs(0)) + + def make_forward(): + @nnx.jit(graph=False, static_argnums=(1, 2)) + def forward(x, mode, scale): + out = model(x) + if mode == 'scale': + return out * scale + return out + return forward + + forward = make_forward() + + x = jnp.ones((4, 2)) + y1 = forward(x, 'scale', 2.0) + y2 = forward(x, 'noscale', 1.0) + np.testing.assert_allclose(y1, y2 * 2.0) + + def test_jit_closure_donates_captured_variables(self): + """Captured nnx.Variable buffers are donated to the jit call.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + model = nnx.Linear(2, 3, rngs=nnx.Rngs(0)) + + @nnx.jit(graph=False) + def forward(x): + return model(x) + + x = jnp.ones((4, 2)) + + old_kernel = model.kernel.get_value() + old_bias = model.bias.get_value() + + _ = forward(x) + + self.assertTrue(old_kernel.is_deleted(), + 'captured kernel buffer was not donated') + self.assertTrue(old_bias.is_deleted(), + 'captured bias buffer was not donated') + + def test_jit_closure_duplicate_error_mentions_captured_args(self): + """Error for duplicate Variable mentions 'captured_args' in the path.""" + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + count = nnx.Variable(jnp.array(0)) + + def make_forward(): + @nnx.jit(graph=False) + def forward(x): + count[...] += 1 + return forward + + forward = make_forward() + + with self.assertRaisesRegex(ValueError, 'captured_args'): + forward(count) + + def test_double_transform_closure(self): + from flax.configurations import config + with config.temp_flip_flag('hijax_variable', True): + x = nnx.Variable(jnp.array(0)) + + @nnx.jit(graph_updates=False) + def f(): + x[...] += 1 + @nnx.jit(graph_updates=False) + def g(): + x[...] += 1 + g() + f() + self.assertEqual(x[...], 2) + f() + self.assertEqual(x[...], 4) + + if __name__ == '__main__': absltest.main()