Skip to content

Commit 244b95b

Browse files
committed
Casting PartitionSepc to Tuple for b/505613482
1 parent 387dc60 commit 244b95b

4 files changed

Lines changed: 9 additions & 9 deletions

File tree

axlearn/common/attention_test.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6433,7 +6433,7 @@ def test_logit_sink_parameter_initialization(self):
64336433
self.assertIn("sink", param_specs)
64346434
sink_spec = param_specs["sink"]
64356435
self.assertEqual(sink_spec.shape, (num_heads,))
6436-
self.assertEqual(sink_spec.mesh_axes, ("model",))
6436+
self.assertEqual(tuple(sink_spec.mesh_axes), ("model",))
64376437
self.assertEqual(sink_spec.weight_decay_scale, 0.0)
64386438

64396439
def test_logit_sink_disabled_by_default(self):

axlearn/common/flash_attention/layer_test.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -526,7 +526,7 @@ def as_partition_spec(pytree: CompositeAttentionBias) -> PartitionSpec:
526526
spec = test_layer._logit_biases_spec(segment_ids) # pylint: disable=protected-access
527527
spec = as_partition_spec(spec)
528528
self.assertIsInstance(spec, PartitionSpec)
529-
self.assertEqual(spec, test_layer.config.mha_dim_to_partition_spec["btnh"][:2])
529+
self.assertEqual(tuple(spec), test_layer.config.mha_dim_to_partition_spec["btnh"][:2])
530530

531531
@parameterized.product(
532532
_TEST_CONFIGS,

axlearn/common/layers_test.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -572,14 +572,14 @@ def test_rms_norm_partition_specs_constraint(self, mock_with_sharding_constraint
572572

573573
# 1. Input tensor constraint.
574574
input_spec = calls[0].args[1]
575-
self.assertEqual(input_spec, ("fsdp", "model", None))
575+
self.assertEqual(tuple(input_spec), ("fsdp", "model", None))
576576
self.assertEqual(calls[0].args[0].shape, (2, 3, dim))
577577
self.assertEqual(calls[0].args[0].dtype, jnp.float32)
578578
np.testing.assert_array_equal(calls[0].args[0], inputs)
579579

580580
# 2. Output tensor constraint.
581581
output_spec = calls[1].args[1]
582-
self.assertEqual(output_spec, ("fsdp", None, None))
582+
self.assertEqual(tuple(output_spec), ("fsdp", None, None))
583583
self.assertEqual(calls[1].args[0].shape, (2, 3, dim))
584584
self.assertEqual(calls[1].args[0].dtype, jnp.float32)
585585
np.testing.assert_array_equal(calls[1].args[0], outputs)
@@ -1462,21 +1462,21 @@ def test_embed_partition_specs_constraint(self, mock_with_sharding_constraint):
14621462

14631463
# 1. Input activation constraint (indices tensor).
14641464
input_spec = calls[0].args[1]
1465-
self.assertEqual(input_spec, ("fsdp", None))
1465+
self.assertEqual(tuple(input_spec), ("fsdp", None))
14661466
self.assertEqual(calls[0].args[0].shape, (3, seq_len))
14671467
self.assertEqual(calls[0].args[0].dtype, jnp.int32)
14681468
np.testing.assert_array_equal(calls[0].args[0], ixs)
14691469

14701470
# 2. Embedding weight constraint.
14711471
weight_spec = calls[1].args[1]
1472-
self.assertEqual(weight_spec, ("model", "fsdp"))
1472+
self.assertEqual(tuple(weight_spec), ("model", "fsdp"))
14731473
self.assertEqual(calls[1].args[0].shape, (num_embeddings, dim))
14741474
self.assertEqual(calls[1].args[0].dtype, jnp.float32)
14751475
np.testing.assert_array_equal(calls[1].args[0], state["weight"])
14761476

14771477
# 3. Output activation constraint (after lookup).
14781478
output_spec = calls[2].args[1]
1479-
self.assertEqual(output_spec, ("fsdp", "model"))
1479+
self.assertEqual(tuple(output_spec), ("fsdp", "model"))
14801480
self.assertEqual(calls[2].args[0].shape, (3, seq_len, dim))
14811481
self.assertEqual(calls[2].args[0].dtype, jnp.float32)
14821482
np.testing.assert_array_equal(calls[2].args[0], actual_embeds)

axlearn/common/lora_test.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -50,8 +50,8 @@ def test_set_param_spec_config(self):
5050
layer_cfg.lora_up.param_partition_spec = [None, "data"]
5151
layer = layer_cfg.instantiate(parent=None)
5252
param_specs = layer.create_parameter_specs_recursively()
53-
self.assertEqual(param_specs["lora_down"]["weight"].mesh_axes, ("data", None))
54-
self.assertEqual(param_specs["lora_up"]["weight"].mesh_axes, (None, "data"))
53+
self.assertEqual(tuple(param_specs["lora_down"]["weight"].mesh_axes), ("data", None))
54+
self.assertEqual(tuple(param_specs["lora_up"]["weight"].mesh_axes), (None, "data"))
5555

5656
def test_forward(self):
5757
input_dim = 2

0 commit comments

Comments
 (0)