diff --git a/src/maxtext/common/common_types.py b/src/maxtext/common/common_types.py index 9dd4aff3dc..051952f764 100644 --- a/src/maxtext/common/common_types.py +++ b/src/maxtext/common/common_types.py @@ -31,6 +31,8 @@ AxisNames = tuple[str, ...] AxisIdxes = tuple[int, ...] +DATA_EMB_BATCH = "activation_embed_and_logits_batch" + BATCH = "activation_batch" BATCH_ATTN = "activation_batch_attn" diff --git a/src/maxtext/layers/attention_op.py b/src/maxtext/layers/attention_op.py index b3f22fc336..0f2c03af95 100644 --- a/src/maxtext/layers/attention_op.py +++ b/src/maxtext/layers/attention_op.py @@ -49,6 +49,7 @@ CACHE_SCALE_SEQUENCE, CACHE_SEQUENCE, Config, + DATA_EMB_BATCH, DECODE_BATCH, DECODE_LENGTH, DECODING_ACTIVE_SEQUENCE_INDICATOR, @@ -747,7 +748,12 @@ def generate_attention_mask( if model_mode == MODEL_MODE_AUTOREGRESSIVE and decoder_segment_ids is not None: mask = decoder_segment_ids[:, None, None, None, :] == DECODING_ACTIVE_SEQUENCE_INDICATOR elif decoder_segment_ids is not None: - mask = decoder_segment_ids[:, :, None] == decoder_segment_ids[:, None, :] + + # With TSP/CP, all-gather prior to broadcast to avoid large all-to-all on broadcasted ids. + key_sharding = self._logical_to_mesh_axes((DATA_EMB_BATCH, None)) + decoder_key_segment_ids = self._maybe_shard_with_pspec(decoder_segment_ids, key_sharding) + + mask = decoder_segment_ids[:, :, None] == decoder_key_segment_ids[:, None, :] mask = mask[:, None, None, :, :] _, q_seq_len, _, _ = query.shape diff --git a/tests/unit/attention_test.py b/tests/unit/attention_test.py index 754c526b2f..2a6c90fa4c 100644 --- a/tests/unit/attention_test.py +++ b/tests/unit/attention_test.py @@ -370,8 +370,16 @@ def test_load_balanced_chunk_window(self): np.testing.assert_array_equal((causal_mask & chunk_mask)[:, :], expected_mask) def test_dot_product_local_mask_uses_segment_positions(self): - config = types.SimpleNamespace(context_parallel_load_balance=True, context_sharding="context") - mesh = types.SimpleNamespace(shape={"context": 4}) + config = types.SimpleNamespace( + context_parallel_load_balance=True, + context_sharding="context", + using_pipeline_parallelism=False, + logical_axis_rules=[["activation_embed_and_logits_batch", ["context"]]], + shard_mode="auto", + debug_sharding=False, + eval_interval=-1, + ) + mesh = Mesh(np.arange(4), ["context"]) seq_len = 16 sliding_window_size = 4 positions = jnp.asarray(attention_op.LoadBalancedCausalMask(shape=(seq_len, seq_len), cp_size=4).q_sequence[None, :]) @@ -406,8 +414,16 @@ def test_dot_product_local_mask_uses_segment_positions(self): np.testing.assert_array_equal(np.asarray(mask == 0.0)[0, 0, 0], expected_mask) def test_dot_product_chunk_mask_uses_segment_positions(self): - config = types.SimpleNamespace(context_parallel_load_balance=True, context_sharding="context") - mesh = types.SimpleNamespace(shape={"context": 4}) + config = types.SimpleNamespace( + context_parallel_load_balance=True, + context_sharding="context", + using_pipeline_parallelism=False, + logical_axis_rules=[["activation_embed_and_logits_batch", ["context"]]], + shard_mode="auto", + debug_sharding=False, + eval_interval=-1, + ) + mesh = Mesh(np.arange(4), ["context"]) seq_len = 16 chunk_size = 4 positions = jnp.asarray(attention_op.LoadBalancedCausalMask(shape=(seq_len, seq_len), cp_size=4).q_sequence[None, :])