Skip to content

Commit a037d08

Browse files
committed
Instead of casting PartitionSpec Tuples to standard Tuples, convert standard
Tuples to PartitionSpecs for comparision statements.
1 parent c3d7ddd commit a037d08

4 files changed

Lines changed: 6 additions & 6 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(tuple(sink_spec.mesh_axes), ("model",))
6436+
self.assertEqual((sink_spec.mesh_axes), jax.P("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(tuple(spec), test_layer.config.mha_dim_to_partition_spec["btnh"][:2])
529+
self.assertEqual(spec, jax.P(test_layer.config.mha_dim_to_partition_spec["btnh"][:2]))
530530

531531
@parameterized.product(
532532
_TEST_CONFIGS,

axlearn/common/layers_test.py

Lines changed: 2 additions & 2 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(tuple(input_spec), ("fsdp", "model", None))
575+
self.assertEqual((input_spec), jax.P("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(tuple(output_spec), ("fsdp", None, None))
582+
self.assertEqual((output_spec), jax.P("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)

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(tuple(param_specs["lora_down"]["weight"].mesh_axes), ("data", None))
54-
self.assertEqual(tuple(param_specs["lora_up"]["weight"].mesh_axes), (None, "data"))
53+
self.assertEqual((param_specs["lora_down"]["weight"].mesh_axes), jax.P("data", None))
54+
self.assertEqual((param_specs["lora_up"]["weight"].mesh_axes), jax.P(None, "data"))
5555

5656
def test_forward(self):
5757
input_dim = 2

0 commit comments

Comments
 (0)