Skip to content

Commit 15f281b

Browse files
lukebaumannGoogle-ML-Automation
authored andcommitted
Implement checkpoint-based elasticity using set-based slice tracking in MaxText.
PiperOrigin-RevId: 948165481
1 parent f835ffb commit 15f281b

2 files changed

Lines changed: 27 additions & 6 deletions

File tree

src/maxtext/utils/elastic_utils.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -211,7 +211,7 @@ def is_scale_up_event(config) -> bool:
211211
if elastic_enabled(config):
212212
ensure_elastic_manager_initialized(config)
213213
assert elastic_manager is not None
214-
return elastic_manager.new_slice_event.is_set()
214+
return bool(elastic_manager.available_inactive_slices)
215215

216216
return False
217217

tests/unit/elastic_utils_test.py

Lines changed: 26 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
import unittest
1818
from unittest.mock import create_autospec, Mock
19+
from absl.testing import parameterized
1920

2021

2122
from 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

451472
if __name__ == "__main__":
452473
unittest.main()

0 commit comments

Comments
 (0)