From 8da934dcbf1021c987b6dc2d1b98a70aa189318d Mon Sep 17 00:00:00 2001 From: maxtext authors Date: Thu, 6 Aug 2026 18:59:45 -0700 Subject: [PATCH] Truncate out_sharding to match tensor ndim when pspec has extra dimensions. PiperOrigin-RevId: 960628527 --- src/maxtext/layers/embeddings.py | 6 +++++- src/maxtext/layers/linears.py | 5 +++++ src/maxtext/layers/normalizations.py | 4 ++++ src/maxtext/utils/sharding.py | 21 ++++++++++++++++++ tests/unit/sharding_nnx_test.py | 32 ++++++++++++++++++++++++++++ 5 files changed, 67 insertions(+), 1 deletion(-) diff --git a/src/maxtext/layers/embeddings.py b/src/maxtext/layers/embeddings.py index 05bc1fa193..bae4514ec9 100644 --- a/src/maxtext/layers/embeddings.py +++ b/src/maxtext/layers/embeddings.py @@ -30,7 +30,7 @@ from maxtext.layers.initializers import Initializer, default_embed_init, variable_to_logically_partitioned from maxtext.utils import max_logging from maxtext.utils import max_utils -from maxtext.utils.sharding import logical_to_mesh_axes, create_sharding +from maxtext.utils.sharding import logical_to_mesh_axes, create_sharding, truncate_out_sharding _MAX_WAVELENGTH = 10_000 @@ -173,6 +173,8 @@ def __call__(self, inputs: Array, model_mode: str = MODEL_MODE_TRAIN) -> Array: out_pspec = logical_to_mesh_axes(output_axis_names, self.mesh, rules=getattr(self.config, "logical_axis_rules", None)) out_sharding = NamedSharding(self.mesh, out_pspec) if self.config.shard_mode == ShardMode.EXPLICIT else None + if out_sharding is not None: + out_sharding = truncate_out_sharding(out_sharding, inputs.ndim + 1) if cfg.use_iota_embed: iota = lax.iota(jnp.int32, self.num_embeddings) @@ -227,6 +229,8 @@ def attend_on_embedding( # out_sharding must be None under auto shard_mode if config.shard_mode != ShardMode.EXPLICIT: out_sharding = None + if out_sharding is not None: + out_sharding = truncate_out_sharding(out_sharding, query.ndim) embedding_table = _maybe_move_embedding_to_device(embedding_table, config) return jnp.dot( query, diff --git a/src/maxtext/layers/linears.py b/src/maxtext/layers/linears.py index 915bc9c008..8e14d6d862 100644 --- a/src/maxtext/layers/linears.py +++ b/src/maxtext/layers/linears.py @@ -41,6 +41,7 @@ from maxtext.utils.sharding import maybe_shard_with_name from maxtext.utils.sharding import get_physical_spec_without_axes from maxtext.utils.sharding import FSDP_MESH_AXES +from maxtext.utils.sharding import truncate_out_sharding def _convert_to_activation_function(fn_or_string: str | Callable[..., Any]) -> Callable[..., Any]: @@ -102,6 +103,10 @@ def _compute_dot_general_nnx( quant_dot_general.lazy_init(inputs, kernel, ((axis, contract_ind), ((), ())), precision=None) return quant_dot_general(inputs, kernel, ((axis, contract_ind), ((), ())), precision=None, mutable=["aqt"]) + if out_sharding is not None: + out_ndim = (inputs.ndim - len(axis)) + (kernel.ndim - len(contract_ind)) + out_sharding = truncate_out_sharding(out_sharding, out_ndim) + return dot_general( inputs, kernel, ((axis, contract_ind), ((), ())), precision=matmul_precision, out_sharding=out_sharding ) diff --git a/src/maxtext/layers/normalizations.py b/src/maxtext/layers/normalizations.py index e98977c60c..2d7ec059ee 100644 --- a/src/maxtext/layers/normalizations.py +++ b/src/maxtext/layers/normalizations.py @@ -28,6 +28,7 @@ from maxtext.layers.initializers import Initializer, variable_to_logically_partitioned from maxtext.utils import max_logging from maxtext.utils import max_utils +from maxtext.utils.sharding import truncate_out_sharding class RMSNorm(nnx.Module): @@ -76,6 +77,9 @@ def __call__(self, x: jnp.ndarray, out_sharding: NamedSharding | None = None) -> if self.shard_mode != ShardMode.EXPLICIT: out_sharding = None + if out_sharding is not None: + out_sharding = truncate_out_sharding(out_sharding, y.ndim) + if not self.with_scale: if out_sharding is not None: y = jax.lax.with_sharding_constraint(y, out_sharding) diff --git a/src/maxtext/utils/sharding.py b/src/maxtext/utils/sharding.py index e02e598a12..5e1983ebf7 100644 --- a/src/maxtext/utils/sharding.py +++ b/src/maxtext/utils/sharding.py @@ -86,6 +86,8 @@ def maybe_shard_with_name( """ if inputs is None: return None + if hasattr(inputs, "ndim"): + named_sharding = truncate_out_sharding(named_sharding, inputs.ndim) if ( isinstance(named_sharding, NamedSharding) and hasattr(inputs, "shape") @@ -348,6 +350,25 @@ def create_sharding(mesh, logical_names, rules=None): return NamedSharding(mesh, logical_to_mesh_axes(logical_names, mesh, rules=rules)) +def truncate_out_sharding(out_sharding, out_ndim: int): + """Truncates out_sharding if tensor ndim is less than out_sharding pspec length.""" + if out_sharding is None: + return None + if isinstance(out_sharding, NamedSharding): + if len(out_sharding.spec) > out_ndim: + return NamedSharding( + out_sharding.mesh, + P(*out_sharding.spec[:out_ndim]), + ) + elif isinstance(out_sharding, P): + if len(out_sharding) > out_ndim: + return P(*out_sharding[:out_ndim]) + elif isinstance(out_sharding, (tuple, list)): + if len(out_sharding) > out_ndim: + return tuple(out_sharding[:out_ndim]) + return out_sharding + + def get_mesh_axes_used_by_tensor_spec(tensor_sharding_spec): """ Extracts the set of mesh axis names that a tensor's PartitionSpec uses. diff --git a/tests/unit/sharding_nnx_test.py b/tests/unit/sharding_nnx_test.py index a073a9d1d3..c9f537b51b 100644 --- a/tests/unit/sharding_nnx_test.py +++ b/tests/unit/sharding_nnx_test.py @@ -408,5 +408,37 @@ def test_removes_size_one_mesh_axes(self): sharding.remove_size_one_mesh_axis = lambda spec, mesh: spec +class TruncateOutShardingTest(unittest.TestCase): + + def setUp(self): + super().setUp() + self.mesh = _create_2d_test_mesh(("data", "model")) + + def test_truncate_out_sharding_none(self): + self.assertIsNone(sharding.truncate_out_sharding(None, 3)) + + def test_truncate_out_sharding_named_sharding(self): + ns = NamedSharding(self.mesh, PartitionSpec(("data", "model"), None, None, None)) + truncated = sharding.truncate_out_sharding(ns, 3) + self.assertIsInstance(truncated, NamedSharding) + self.assertEqual(truncated.mesh, self.mesh) + self.assertEqual(truncated.spec, PartitionSpec(("data", "model"), None, None)) + + def test_truncate_out_sharding_named_sharding_no_op(self): + ns = NamedSharding(self.mesh, PartitionSpec(("data", "model"), None)) + truncated = sharding.truncate_out_sharding(ns, 3) + self.assertIs(truncated, ns) + + def test_truncate_out_sharding_partition_spec(self): + pspec = PartitionSpec("data", "model", None, None) + truncated = sharding.truncate_out_sharding(pspec, 2) + self.assertEqual(truncated, PartitionSpec("data", "model")) + + def test_truncate_out_sharding_tuple(self): + spec = ("data", "model", None, None) + truncated = sharding.truncate_out_sharding(spec, 2) + self.assertEqual(truncated, ("data", "model")) + + if __name__ == "__main__": unittest.main()