Repository navigation
Out sharding for modules initialized with JIT is incorrect #5127
Description
Activity
I've also tried couple of methods to specify out_sharding to nnx.jit not neither of them worked for me:
import jax import flax.nnx as nnx import jax import jax.numpy as jnp import flax.nnx as nnx import traceback mesh1 = jax.make_mesh((2, 4), ("a", "b")) rules1 = (("A", "a"), ("B", "b")) mesh2 = jax.make_mesh((2, 2, 2), ("x", "y", "z")) rules2 = (("X", "x"), ("Y", "y"), ("Z", "z")) mesh3 = jax.make_mesh((8,), ("c",)) rules3 = (("C", "c"),) mesh_data = jax.make_mesh((4, 2), ("data", "context")) def eval_shape_with_sharding(f, *args, **kwargs): # Currently flax's nnx.eval_shape does not propagate sharding information. # Issue to track: https://github.com/google/flax/issues/5110 module = nnx.eval_shape(f, *args, **kwargs) state = nnx.state(module) pspec = nnx.spmd.get_partition_spec(state) def wrap_with_sharding(var: nnx.Variable, var_pspec: nnx.Variable) -> nnx.Variable: value = var.get_value() if not isinstance(value, jax.ShapeDtypeStruct | jax.Array): # var.value may be MaskedNode when training subset of parameters return var new_var = var.copy() var_mesh = var.get_metadata().get("mesh", None) if var_mesh is not None: new_var.set_value(jax.ShapeDtypeStruct( shape=value.shape, dtype=value.dtype, sharding=jax.sharding.NamedSharding( mesh=var_mesh, spec=var_pspec.get_value(), ), )) return new_var state_with_sharding = jax.tree.map( lambda t, spec: wrap_with_sharding(t, spec), state, pspec, is_leaf=lambda x: isinstance(x, nnx.Variable), ) nnx.update(module, state_with_sharding) return module class Model(nnx.Module): def __init__(self): self.small_linear1 = nnx.Param( jnp.ones((16, 16)), sharding=("A", "B"), mesh=mesh1, sharding_rules=rules1, ) self.small_linear2 = nnx.Param( jnp.ones((16, 16, 16)), sharding=("X", "Y", "Z"), mesh=mesh2, sharding_rules=rules2, ) self.small_linear3 = nnx.Param( jnp.ones((16, 16)), sharding=("C",), mesh=mesh3, sharding_rules=rules3, ) def init_model_no_jit(): return Model() abstract_model = eval_shape_with_sharding(init_model_no_jit) with mesh_data: @nnx.jit(out_shardings=jax.tree.map(lambda t: t.sharding, abstract_model)) def init_model_nnx_jit(): model = init_model_no_jit() return model print("-" * 100) try: model_nnx_jit = init_model_nnx_jit() except Exception as e: print("nnx.jit with out_sharding from state failed with exception:") traceback.print_exc() @nnx.jit(out_shardings=jax.tree.map(lambda t: t.sharding, nnx.state(abstract_model))) def init_model_nnx_jit_out_state(): model = init_model_no_jit() return model print("-" * 100) try: model_nnx_jit = init_model_nnx_jit_out_state() except Exception as e: print("nnx.jit with out_sharding from state failed with exception:") traceback.print_exc() @nnx.jit(out_shardings=jax.tree.map(lambda t: t.sharding, nnx.state(abstract_model))) def init_model_nnx_jit_out_state_return_state(): model = init_model_no_jit() return nnx.state(model) print("-" * 100) try: model_nnx_jit = init_model_nnx_jit_out_state_return_state() except Exception as e: print("nnx.jit with out_sharding from state returning state failed with exception:") traceback.print_exc()
output:
nnx.jit failed with exception: Traceback (most recent call last): File "/papyrax/test_sharding_on_init.py", line 92, in <module> model_nnx_jit = init_model_nnx_jit() ^^^^^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/transforms/compilation.py", line 474, in __call__ pure_args_out, pure_kwargs_out, pure_out = self.jitted_fn( ^^^^^^^^^^^^^^^ ValueError: pytree structure error: different types at key path pjit out_shardings[2] At that key path, the prefix pytree pjit out_shardings has a subtree of type <class '__main__.Model'> but at the same key path the full pytree has a subtree of different type <class 'flax.nnx.extract.NodeStates'>. -------------------- For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these. ---------------------------------------------------------------------------------------------------- nnx.jit with out_sharding from state failed with exception: Traceback (most recent call last): File "/papyrax/test_sharding_on_init.py", line 104, in <module> model_nnx_jit = init_model_nnx_jit_out_state() ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/transforms/compilation.py", line 474, in __call__ pure_args_out, pure_kwargs_out, pure_out = self.jitted_fn( ^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/transforms/compilation.py", line 138, in __call__ pure_args_out, pure_kwargs_out, pure_out = extract.to_tree( ^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/extract.py", line 225, in to_tree leaf_prefixes = broadcast_prefix( ^^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/extract.py", line 126, in broadcast_prefix jax.tree.map( File "/usr/local/lib/python3.11/dist-packages/jax/_src/tree.py", line 155, in map return tree_util.tree_map(f, tree, *rest, is_leaf=is_leaf) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ ValueError: Custom node type mismatch: expected type: <class 'flax.nnx.statelib.State'>, value: Model( # Param: 4,608 (18.4 KB) small_linear1=Param( # 256 (1.0 KB) value=Array(shape=(16, 16), dtype=dtype('float32')), mesh=Mesh(axis_sizes=(2, 4), axis_names=('a', 'b'), axis_types=(Auto, Auto)), sharding_rules=(('A', 'a'), ('B', 'b')), sharding_names=('A', 'B') ), small_linear2=Param( # 4,096 (16.4 KB) value=Array(shape=(16, 16, 16), dtype=dtype('float32')), mesh=Mesh(axis_sizes=(2, 2, 2), axis_names=('x', 'y', 'z'), axis_types=(Auto, Auto, Auto)), sharding_rules=(('X', 'x'), ('Y', 'y'), ('Z', 'z')), sharding_names=('X', 'Y', 'Z') ), small_linear3=Param( # 256 (1.0 KB) value=Array(shape=(16, 16), dtype=dtype('float32')), mesh=Mesh(axis_sizes=(8,), axis_names=('c',), axis_types=(Auto,)), sharding_rules=(('C', 'c'),), sharding_names=('C',) ) ). -------------------- For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these. ---------------------------------------------------------------------------------------------------- nnx.jit with out_sharding from state returning state failed with exception: Traceback (most recent call last): File "/papyrax/test_sharding_on_init.py", line 116, in <module> model_nnx_jit = init_model_nnx_jit_out_state_return_state() ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/transforms/compilation.py", line 474, in __call__ pure_args_out, pure_kwargs_out, pure_out = self.jitted_fn( ^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/transforms/compilation.py", line 138, in __call__ pure_args_out, pure_kwargs_out, pure_out = extract.to_tree( ^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/extract.py", line 248, in to_tree check_consistent_aliasing( File "/usr/local/lib/python3.11/dist-packages/flax/nnx/extract.py", line 86, in check_consistent_aliasing unique_prefixes = {prefix for _, prefix in paths_prefixes} ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/usr/local/lib/python3.11/dist-packages/flax/nnx/extract.py", line 86, in <setcomp> unique_prefixes = {prefix for _, prefix in paths_prefixes} ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ TypeError: unhashable type: 'Param' -------------------- For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these.Only option that worked for me is to use jax.jit for init which would require additional machinery with split/merge that I would like to avoid if possible.
with mesh_data: @jax.jit(out_shardings=jax.tree.map(lambda t: t.sharding, nnx.state(abstract_model))) def init_model_nnx_jit(): model = init_model_no_jit() return nnx.state(model) def _print_t_shading(key, t): print(f"Key: {'.'.join(map(str, key))}, shape: {t.shape}, sharding: {t.sharding}") jitted_model = init_model_nnx_jit() print("Sharding with JAX.JIT:") jax.tree.map_with_path(_print_t_shading, jitted_model)
Key: ['small_linear1']..value, shape: (16, 16), sharding: NamedSharding(mesh=Mesh('a': 2, 'b': 4, axis_types=(Auto, Auto)), spec=PartitionSpec('a', 'b'), memory_kind=device) Key: ['small_linear2']..value, shape: (16, 16, 16), sharding: NamedSharding(mesh=Mesh('x': 2, 'y': 2, 'z': 2, axis_types=(Auto, Auto, Auto)), spec=PartitionSpec('x', 'y', 'z'), memory_kind=device) Key: ['small_linear3']..value, shape: (16, 16), sharding: NamedSharding(mesh=Mesh('c': 8, axis_types=(Auto,)), spec=PartitionSpec('c',), memory_kind=device)@qGentry I haven't checked yet the code in details, but
with mesh_data:code is not good. In jax they recommend using new mesh API:with jax.set_mesh(mesh):instead.Here's a slightly smaller reproduction:
mesh1 = jax.make_mesh((2, 4), ("a", "b")) mesh2 = jax.make_mesh((2, 2, 2), ("x", "y", "z")) class Model(nnx.Module): @nnx.jit def __init__(self): self.small_linear1 = nnx.Param( jnp.ones((16, 16)), sharding=("a", "b"), mesh=mesh1, ) self.small_linear2 = nnx.Param( jnp.ones((16, 16, 16)), sharding=("x", "y", "z"), mesh=mesh2, ) def print_sharding(key, t): print(f"Key: {'.'.join(map(str, key))}, shape: {t.shape}, sharding: {t.sharding}") model = Model() jax.tree.map_with_path(print_sharding, model)
As before, with
nnx.jit,small_linear2is sharded as a GSPMDSharding rather than a NamedSharding.When you pass a sharding to
nnx.Param, flax currently reshards the input as follows:def do_shard(a, sharding): with jax.disable_jit(False): return jax.jit(lambda x: x, out_shardings=sharding)(a)
With is in mind, here is an even smaller reproduction:
@jax.jit def jax_model(): a = do_shard(jnp.ones((16, 16)),NamedSharding(mesh1, P('a', 'b'))) b = do_shard(jnp.ones((16, 16, 16)), NamedSharding(mesh2, P('x', 'y', 'z'))) return (a,b) def print_t_shading(key, t): print(f"Key: {'.'.join(map(str, key))}, shape: {t.shape}, sharding: {t.sharding}") model = jax_model() jax.tree.map_with_path(print_t_shading, model)
So what's breaking here is this
do_shardfunction.@qGentry Using multiple different meshes within a
jax.jitcurrently just isn't possible in jax itself. In my example above, jitting needs to resolve an output sharding for the pytree(a,b)that uses a single mesh. Does this make sense?So I believe you'll need to refactor your code to use a single mesh within your compilation context.
Hi @samanklesaria, thanks for looking into it.
JAX works properly with using multiple meshes in compilation context, we've been using this approach for quite a while now (from JAX 0.4.31 to modern 0.8.*). But I agree that this is mostly jax.jit-related issues, in a sense that without specifying out_shardings in jax.jit, JAX tries to match GSPMD sharding to single mesh (although it is weird that it is chosen based on ordering of the meshes). Generally, I was able to avoid this issue by actually specifying out_sharding using this approach #5127 (comment)
Hey folks, me again.
I've recently faced the following problem when initializing the model with multiple meshes. Basically, output sharding from jitted init_fn returns completely random sharding instead of sticking to specified ones. Also seems like output tensors's mesh actually depends on ordering of the flattened tree. Check out this repro script:
output:
No-JIT version, on the other hand, works correctly (but as one may imagine is not suitable for large-scale init).