Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion src/maxtext/layers/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
5 changes: 5 additions & 0 deletions src/maxtext/layers/linears.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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
)
Expand Down
4 changes: 4 additions & 0 deletions src/maxtext/layers/normalizations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Expand Down
21 changes: 21 additions & 0 deletions src/maxtext/utils/sharding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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.
Expand Down
32 changes: 32 additions & 0 deletions tests/unit/sharding_nnx_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading