1616
1717import unittest
1818from unittest .mock import create_autospec , Mock
19+ from absl .testing import parameterized
1920
2021
2122from maxtext .utils import elastic_utils
@@ -46,7 +47,7 @@ def __init__(self):
4647 self .elastic_min_slice_count = 1
4748
4849
49- class ElasticUtilsTest (unittest .TestCase ):
50+ class ElasticUtilsTest (parameterized .TestCase ):
5051 """Unit tests for Elastic Training utility functions."""
5152
5253 def setUp (self ):
@@ -66,7 +67,7 @@ def setUp(self):
6667 self .fake_logging = create_autospec (self .original_max_logging )
6768 self .fake_jax = create_autospec (self .original_jax )
6869 self .fake_manager = create_autospec (self .original_manager_class , instance = True )
69- self .fake_manager .new_slice_event = Mock ()
70+ self .fake_manager .available_inactive_slices = set ()
7071
7172 # Configure default behaviors if needed
7273 self .fake_pathwaysutils .is_pathways_backend_used .return_value = True
@@ -314,7 +315,7 @@ def wait_until_finished(self):
314315 cm = FakeCheckpointManager ()
315316
316317 elastic_utils .elastic_manager = self .fake_manager
317- self .fake_manager .new_slice_event . is_set . return_value = True
318+ self .fake_manager .available_inactive_slices = { 1 }
318319
319320 with self .assertRaises (ScaleUpSignalError ):
320321 elastic_utils .maybe_elastic_scale_up (config , cm )
@@ -359,7 +360,7 @@ def test_elastic_retry_pre_callback_forwarded(self):
359360 def test_record_elastic_event_start (self ):
360361 """Tests recording an elastic slice down start."""
361362 elastic_utils .elastic_manager = self .fake_manager
362- self .fake_manager .new_slice_event . is_set . return_value = False
363+ self .fake_manager .available_inactive_slices = set ()
363364 fake_recorder = Mock ()
364365 config = FakeConfig ()
365366
@@ -373,7 +374,7 @@ def test_record_elastic_event_start(self):
373374 def test_record_elastic_event_start_scale_up (self ):
374375 """Tests recording an elastic slice scale up start."""
375376 elastic_utils .elastic_manager = self .fake_manager
376- self .fake_manager .new_slice_event . is_set . return_value = True
377+ self .fake_manager .available_inactive_slices = { 1 }
377378 fake_recorder = Mock ()
378379 config = FakeConfig ()
379380
@@ -447,6 +448,26 @@ def __setattr__(self, name, value):
447448 elastic_utils .ensure_elastic_manager_initialized (config )
448449 self .assertEqual (elastic_utils .elastic_manager , self .fake_manager )
449450
451+ @parameterized .parameters (
452+ # Positive cases
453+ ({1 }, True ),
454+ ({0 }, True ),
455+ ({1 , 2 }, True ),
456+ ({0 , 3 , 6 }, True ),
457+ ({10 , 25 }, True ),
458+ # Negative cases
459+ (set (), False ),
460+ )
461+ def test_is_scale_up_event_with_set (self , available_inactive_slices , expected ):
462+ config = FakeConfig ()
463+ config .elastic_enabled = True
464+ elastic_utils .elastic_manager = self .fake_manager
465+
466+ self .fake_manager .available_inactive_slices = available_inactive_slices
467+ self .assertEqual (elastic_utils .is_scale_up_event (config ), expected )
468+
469+
470+
450471
451472if __name__ == "__main__" :
452473 unittest .main ()
0 commit comments