@@ -70,12 +70,6 @@ def setUp(self):
7070
7171 # Configure default behaviors if needed
7272 self .fake_pathwaysutils .is_pathways_backend_used .return_value = True
73- self .fake_pathwaysutils .elastic = Mock ()
74- self .fake_pathwaysutils .elastic .get_active_slice_indices .return_value = [0 , 1 ]
75- self .fake_pathwaysutils .elastic .get_slice_to_devices .return_value = {
76- 0 : [FakeDevice ()],
77- 1 : [FakeDevice ()],
78- }
7973 self .fake_jax .process_index .return_value = 0
8074
8175 # Inject fakes into elastic_utils namespace
@@ -103,28 +97,9 @@ def tearDown(self):
10397 elastic_utils .pending_elastic_event_type = None
10498 super ().tearDown ()
10599
106- def test_record_slice_state (self ):
107- elastic_utils .elastic_manager = self .fake_manager
108- self .fake_manager .active_slice_indices = {0 }
109- self .fake_manager .slice_to_devices = {0 : [FakeDevice ()], 1 : [FakeDevice ()]}
110- self .fake_pathwaysutils .elastic .get_active_slice_indices .return_value = {0 , 1 }
111-
112- fake_recorder = Mock ()
113- fake_recorder .record_elastic_slice_counts = Mock ()
114-
115- elastic_utils .record_slice_state (fake_recorder )
116-
117- fake_recorder .record_elastic_slice_counts .assert_called_once_with (available_slices = 2 , active_slices = 1 , total_slices = 2 )
118-
119- fake_recorder .record_elastic_slice_counts .reset_mock ()
120- elastic_utils .record_slice_state (fake_recorder , active_slices_override = 0 )
121- fake_recorder .record_elastic_slice_counts .assert_called_once_with (available_slices = 2 , active_slices = 0 , total_slices = 2 )
122-
123100 def test_elastic_enabled (self ):
124101 config = FakeConfig ()
125102 self .fake_pathwaysutils .is_pathways_backend_used .return_value = True
126- self .fake_pathwaysutils .elastic = Mock ()
127- self .fake_pathwaysutils .elastic .get_active_slice_indices .return_value = [0 , 1 ]
128103 config .elastic_enabled = True
129104 self .assertTrue (elastic_utils .elastic_enabled (config ))
130105
@@ -169,8 +144,6 @@ def test_live_devices_no_pathways(self):
169144 def test_live_devices_pathways (self ):
170145 """Tests live_devices when pathways is used."""
171146 self .fake_pathwaysutils .is_pathways_backend_used .return_value = True
172- self .fake_pathwaysutils .elastic = Mock ()
173- self .fake_pathwaysutils .elastic .get_active_slice_indices .return_value = [0 , 1 ]
174147 device0 = FakeDevice (slice_index = 0 )
175148 device1 = FakeDevice (slice_index = 1 )
176149 self .fake_jax .devices .return_value = [device0 , device1 ]
@@ -195,8 +168,6 @@ def test_live_devices_disabled(self):
195168 def test_elastic_retry_disabled (self ):
196169 """Tests elastic_retry when disabled but pathways is used."""
197170 self .fake_pathwaysutils .is_pathways_backend_used .return_value = True
198- self .fake_pathwaysutils .elastic = Mock ()
199- self .fake_pathwaysutils .elastic .get_active_slice_indices .return_value = [0 , 1 ]
200171 config = FakeConfig ()
201172 config .elastic_enabled = False
202173 msg = (
@@ -348,28 +319,28 @@ def test_record_elastic_event_start(self):
348319 """Tests recording an elastic slice down start."""
349320 elastic_utils .elastic_manager = self .fake_manager
350321 self .fake_manager .new_slice_event .is_set .return_value = False
351- self .fake_manager .slice_to_devices = {0 : [FakeDevice ()], 1 : [FakeDevice ()]}
352322 fake_recorder = Mock ()
353323 config = FakeConfig ()
354324
355325 elastic_utils .record_elastic_event_start (fake_recorder , config )
356326
357- fake_recorder .record_elastic_wait_start_time .assert_called_once_with (event_type = "elastic_slice_down" )
358- fake_recorder .record_elastic_slice_counts .assert_called_once ()
327+ fake_recorder .record_custom_badput_event_start_time .assert_called_once_with (
328+ custom_badput_event_type = "elastic_slice_down"
329+ )
359330 self .assertEqual (elastic_utils .pending_elastic_event_type , "elastic_slice_down" )
360331
361332 def test_record_elastic_event_start_scale_up (self ):
362333 """Tests recording an elastic slice scale up start."""
363334 elastic_utils .elastic_manager = self .fake_manager
364335 self .fake_manager .new_slice_event .is_set .return_value = True
365- self .fake_manager .slice_to_devices = {0 : [FakeDevice ()], 1 : [FakeDevice ()]}
366336 fake_recorder = Mock ()
367337 config = FakeConfig ()
368338
369339 elastic_utils .record_elastic_event_start (fake_recorder , config )
370340
371- fake_recorder .record_elastic_wait_start_time .assert_called_once_with (event_type = "elastic_scale_up" )
372- fake_recorder .record_elastic_slice_counts .assert_called_once ()
341+ fake_recorder .record_custom_badput_event_start_time .assert_called_once_with (
342+ custom_badput_event_type = "elastic_scale_up"
343+ )
373344
374345 def test_record_elastic_wait_end_and_reinit_start_noop_on_first_attempt (self ):
375346 """Tests recording elastic event end and elastic reinit start."""
@@ -378,39 +349,36 @@ def test_record_elastic_wait_end_and_reinit_start_noop_on_first_attempt(self):
378349
379350 elastic_utils .record_elastic_wait_end_and_reinit_start (fake_recorder )
380351
381- fake_recorder .record_elastic_wait_end_time .assert_not_called ()
382- fake_recorder .record_elastic_reinit_start_time .assert_not_called ()
352+ fake_recorder .record_custom_badput_event_end_time .assert_not_called ()
353+ fake_recorder .record_custom_badput_event_start_time .assert_not_called ()
383354 self .assertIsNone (elastic_utils .pending_reinit_recorder )
384355
385356 def test_record_elastic_wait_end_and_reinit_start (self ):
386357 """Test recording end of slice down and start of reinit."""
387358 elastic_utils .pending_elastic_event_type = "elastic_slice_down"
388- elastic_utils .elastic_manager = self .fake_manager
389- self .fake_manager .active_slice_indices = {0 }
390- self .fake_manager .slice_to_devices = {0 : [FakeDevice ()], 1 : [FakeDevice ()]}
391- elastic_utils .pending_elastic_event_type = "elastic_slice_down"
392359 fake_recorder = Mock ()
393360
394361 elastic_utils .record_elastic_wait_end_and_reinit_start (fake_recorder )
395362
396- fake_recorder .record_elastic_wait_end_time .assert_called_once_with (event_type = "elastic_slice_down" )
397- fake_recorder .record_elastic_reinit_start_time .assert_called_once ()
398- fake_recorder .record_elastic_slice_counts .assert_called_once ()
363+ fake_recorder .record_custom_badput_event_end_time .assert_called_once_with (
364+ custom_badput_event_type = "elastic_slice_down"
365+ )
366+ fake_recorder .record_custom_badput_event_start_time .assert_called_once_with (
367+ custom_badput_event_type = "elastic_reinitialization"
368+ )
399369 self .assertIs (elastic_utils .pending_reinit_recorder , fake_recorder )
400370 self .assertIsNone (elastic_utils .pending_elastic_event_type )
401371
402372 def test_record_elastic_reinit_end (self ):
403373 """Tests recording end of elastic reinit."""
404374 fake_recorder = Mock ()
405375 elastic_utils .pending_reinit_recorder = fake_recorder
406- elastic_utils .elastic_manager = self .fake_manager
407- self .fake_manager .active_slice_indices = {0 }
408- self .fake_manager .slice_to_devices = {0 : [FakeDevice ()], 1 : [FakeDevice ()]}
409376
410377 elastic_utils .record_elastic_reinit_end ()
411378
412- fake_recorder .record_elastic_reinit_end_time .assert_called_once ()
413- fake_recorder .record_elastic_slice_counts .assert_called_once ()
379+ fake_recorder .record_custom_badput_event_end_time .assert_called_once_with (
380+ custom_badput_event_type = "elastic_reinitialization"
381+ )
414382 self .assertIsNone (elastic_utils .pending_reinit_recorder )
415383
416384 def test_record_elastic_reinit_end_on_cold_start (self ):
@@ -433,8 +401,6 @@ def __setattr__(self, name, value):
433401
434402 config = ReadOnlyConfig ()
435403 self .fake_pathwaysutils .is_pathways_backend_used .return_value = True
436- self .fake_pathwaysutils .elastic = Mock ()
437- self .fake_pathwaysutils .elastic .get_active_slice_indices .return_value = [0 , 1 ]
438404
439405 # Should not raise ValueError
440406 elastic_utils .ensure_elastic_manager_initialized (config )
0 commit comments