Skip to content

Out sharding for modules initialized with JIT is incorrect #5127

Description

@qGentry

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:

import jax
import flax.nnx as nnx

import jax
import jax.numpy as jnp
import flax.nnx as nnx


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"))


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()


@nnx.jit
def init_model_nnx_jit():
    model = init_model_no_jit()
    return model


with mesh_data:
    model_nnx_jit = init_model_nnx_jit()
    model_no_jit = init_model_no_jit()

    def _print_t_shading(key, t):
        print(f"Key: {'.'.join(map(str, key))}, shape: {t.shape}, sharding: {t.sharding}")

    print("\nSharding without JIT:")
    jax.tree.map_with_path(_print_t_shading, model_no_jit)

    print("Sharding with NNX.JIT:")
    jax.tree.map_with_path(_print_t_shading, model_nnx_jit)

output:

Sharding without JIT:
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)
Sharding with NNX.JIT:
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: GSPMDSharding({devices=[2,2,2]<=[8]}, memory_kind=device)
Key: .small_linear3..value, shape: (16, 16), sharding: NamedSharding(mesh=Mesh('a': 2, 'b': 4, axis_types=(Auto, Auto)), spec=PartitionSpec(('a', 'b'),), memory_kind=device)

No-JIT version, on the other hand, works correctly (but as one may imagine is not suitable for large-scale init).

Activity

  1. qGentry commented on Dec 8, 2025

    @qGentry
    Author

    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.
    
  2. qGentry commented on Dec 8, 2025

    @qGentry
    Author

    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)
    
  3. vfdev-5 commented on Dec 9, 2025

    @vfdev-5
    Collaborator

    @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.

  4. samanklesaria commented on Dec 11, 2025

    @samanklesaria
    Collaborator

    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_linear2 is sharded as a GSPMDSharding rather than a NamedSharding.

  5. samanklesaria commented on Dec 11, 2025

    @samanklesaria
    Collaborator

    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_shard function.

  6. samanklesaria commented on Dec 11, 2025

    @samanklesaria
    Collaborator

    @qGentry Using multiple different meshes within a jax.jit currently 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.

  7. qGentry commented on Dec 17, 2025

    @qGentry
    Author

    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)

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions