diff --git a/flax/nnx/bridge/wrappers.py b/flax/nnx/bridge/wrappers.py index a68dfa141..b3117ad23 100644 --- a/flax/nnx/bridge/wrappers.py +++ b/flax/nnx/bridge/wrappers.py @@ -323,6 +323,21 @@ class ToLinen(linen.Module): skip_rng: bool = False metadata_fn: tp.Callable[[variablelib.Variable], tp.Any] | None = bv.to_linen_var + def __post_init__(self): + super_post_init = getattr(super(), '__post_init__', None) + if super_post_init is not None: + super_post_init() + if not isinstance(self.args, tuple): + object.__setattr__( + self, 'args', tuple(self.args) if self.args is not None else () + ) + if not isinstance(self.kwargs, FrozenDict): + object.__setattr__( + self, + 'kwargs', + FrozenDict(self.kwargs) if self.kwargs is not None else FrozenDict({}), + ) + @linen.compact def __call__( self, *args, nnx_method: tp.Callable[..., Any] | str | None = None, **kwargs diff --git a/tests/nnx/bridge/wrappers_test.py b/tests/nnx/bridge/wrappers_test.py index 6d0a38a77..7ef3de1de 100644 --- a/tests/nnx/bridge/wrappers_test.py +++ b/tests/nnx/bridge/wrappers_test.py @@ -325,6 +325,38 @@ def __call__(self): assert updates['Count']['count'] == 1 _ = model.apply(variables | updates) + def test_to_linen_hashable(self): + # Default ToLinen + model1 = bridge.ToLinen(nnx.Linear, args=(32, 64)) + hash1 = hash(model1) + self.assertIsInstance(hash1, int) + + # ToLinen with mutable dict kwargs should be auto-frozen and hashable + model2 = bridge.ToLinen( + nnx.Linear, args=(32, 64), kwargs={'use_bias': False} + ) + hash2 = hash(model2) + self.assertIsInstance(hash2, int) + self.assertIsInstance(model2.kwargs, flax.core.FrozenDict) + + # ToLinen with list args should be converted to tuple and hashable + model3 = bridge.ToLinen( + nnx.Linear, args=[32, 64], kwargs={'use_bias': False} + ) + hash3 = hash(model3) + self.assertEqual(hash2, hash3) + self.assertIsInstance(model3.args, tuple) + + # Usable in jax.jit as static argument + def apply_fn(module, variables, x): + return module.apply(variables, x) + + jitted_apply = jax.jit(apply_fn, static_argnums=(0,)) + x = jax.numpy.ones((1, 32)) + variables = model2.init(jax.random.key(0), x) + out = jitted_apply(model2, variables, x) + self.assertEqual(out.shape, (1, 64)) + def test_to_linen_method_call(self): class Foo(nn.Module): def setup(self):