Skip to content

Commit 0cec2ab

Browse files
Merge pull request #4350 from AI-Hypercomputer:sujinesh/pathways-mtc-active-elastic-devices
PiperOrigin-RevId: 944709894
2 parents 909fd01 + 422ed52 commit 0cec2ab

4 files changed

Lines changed: 111 additions & 2 deletions

File tree

src/maxtext/utils/elastic_utils.py

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

1717
import functools
1818
from collections import Counter
19+
from types import SimpleNamespace
1920

2021
import jax
2122
from maxtext.utils import gcs_utils
@@ -226,3 +227,34 @@ def maybe_elastic_scale_up(config, checkpoint_manager):
226227
checkpoint_manager.wait_until_finished()
227228
max_logging.log("Checkpoint save completed. Interrupting")
228229
raise manager.ScaleUpSignalError()
230+
231+
232+
def single_controller_mtc_init_kwargs(raw_keys):
233+
"""Returns topology kwargs for single-controller MTC initialization."""
234+
kwargs = {
235+
"data_parallelism": raw_keys["mtc_data_parallelism"],
236+
"num_slices": raw_keys["num_slices"],
237+
}
238+
if not raw_keys.get("elastic_enabled", False):
239+
return kwargs
240+
241+
config = SimpleNamespace(**raw_keys)
242+
if not should_use_elastic(config):
243+
return kwargs
244+
245+
active_devices = tuple(live_devices(config))
246+
active_slice_indices = {getattr(device, "slice_index", 0) for device in active_devices if device is not None}
247+
if not active_devices or not active_slice_indices:
248+
raise ValueError("Elastic single-controller MTC initialization found no active devices.")
249+
250+
kwargs["devices"] = active_devices
251+
kwargs["num_slices"] = len(active_slice_indices)
252+
if not kwargs["data_parallelism"]:
253+
kwargs["data_parallelism"] = kwargs["num_slices"]
254+
max_logging.log(
255+
"Using active elastic devices for single-controller MTC initialization: "
256+
f"active_num_slices={kwargs['num_slices']}, "
257+
f"active_device_count={len(active_devices)}, "
258+
f"configured_num_slices={raw_keys['num_slices']}."
259+
)
260+
return kwargs

src/maxtext/utils/max_utils.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -248,14 +248,14 @@ def maybe_initialize_jax_distributed_system(raw_keys):
248248
max_logging.log("Skipping jax distributed system since its not needed for single controller.")
249249
if raw_keys["enable_multi_tier_checkpointing"]:
250250
max_logging.log("Initializing multi-tier checkpointing for single controller...")
251+
mtc_init_kwargs = elastic_utils.single_controller_mtc_init_kwargs(raw_keys)
251252
initialize_multi_tier_checkpointing(
252253
local_checkpoint_directory=raw_keys["local_checkpoint_directory"],
253254
backup_interval_minutes=raw_keys["multi_tier_checkpointing_backup_interval_minutes"],
254255
run_name=raw_keys["run_name"],
255256
jax_initialization_timeout_seconds=raw_keys["jax_distributed_initialization_timeout"],
256-
data_parallelism=raw_keys["mtc_data_parallelism"],
257-
num_slices=raw_keys["num_slices"],
258257
use_colocated_python=True,
258+
**mtc_init_kwargs,
259259
)
260260
return
261261
if jax.distributed.is_initialized():

tests/unit/elastic_utils_test.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -245,6 +245,45 @@ def test_live_slice_indices(self):
245245
indices = elastic_utils.live_slice_indices(config)
246246
self.assertEqual(indices, {0, 1})
247247

248+
def _base_mtc_keys(self, **overrides):
249+
keys = {
250+
"elastic_enabled": True,
251+
"mtc_data_parallelism": 1,
252+
"num_slices": 2,
253+
}
254+
keys.update(overrides)
255+
return keys
256+
257+
def test_single_controller_mtc_init_kwargs_uses_active_elastic_devices(self):
258+
self.fake_pathwaysutils.is_pathways_backend_used.return_value = True
259+
device0 = FakeDevice(slice_index=0)
260+
device1 = FakeDevice(slice_index=1)
261+
self.fake_jax.devices.return_value = [device0, device1]
262+
self.fake_manager.active_slice_indices = {0}
263+
264+
kwargs = elastic_utils.single_controller_mtc_init_kwargs(self._base_mtc_keys(mtc_data_parallelism=0))
265+
266+
self.assertEqual(kwargs["devices"], (device0,))
267+
self.assertEqual(kwargs["num_slices"], 1)
268+
self.assertEqual(kwargs["data_parallelism"], 1)
269+
270+
def test_single_controller_mtc_init_kwargs_raises_if_empty(self):
271+
self.fake_pathwaysutils.is_pathways_backend_used.return_value = True
272+
device0 = FakeDevice(slice_index=0)
273+
self.fake_jax.devices.return_value = [device0]
274+
self.fake_manager.active_slice_indices = {1}
275+
276+
with self.assertRaisesRegex(ValueError, "Elastic single-controller MTC initialization found no active devices."):
277+
elastic_utils.single_controller_mtc_init_kwargs(self._base_mtc_keys())
278+
279+
def test_single_controller_mtc_init_kwargs_non_elastic(self):
280+
kwargs = elastic_utils.single_controller_mtc_init_kwargs(
281+
self._base_mtc_keys(elastic_enabled=False, mtc_data_parallelism=3, num_slices=4)
282+
)
283+
284+
self.assertEqual(kwargs, {"data_parallelism": 3, "num_slices": 4})
285+
self.assertIsNone(elastic_utils.elastic_manager)
286+
248287
def test_get_devices_per_host(self):
249288
device0 = FakeDevice(slice_index=0, process_index=0, task_id=0)
250289
device1 = FakeDevice(slice_index=0, process_index=0, task_id=0)

tests/unit/max_utils_test.py

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -492,6 +492,44 @@ def test_single_controller_multi_tier_checkpointing_uses_colocated_python(self,
492492
use_colocated_python=True,
493493
)
494494

495+
@mock.patch("maxtext.utils.max_utils.elastic_utils.single_controller_mtc_init_kwargs")
496+
@mock.patch("maxtext.utils.max_utils.initialize_multi_tier_checkpointing")
497+
@mock.patch("jax.distributed.initialize")
498+
def test_single_controller_multi_tier_checkpointing_uses_elastic_utils_kwargs(
499+
self, mock_init, mock_mtc, mock_mtc_init_kwargs
500+
):
501+
active_devices = (
502+
mock.Mock(slice_index=0),
503+
mock.Mock(slice_index=0),
504+
)
505+
mock_mtc_init_kwargs.return_value = {
506+
"data_parallelism": 1,
507+
"num_slices": 1,
508+
"devices": active_devices,
509+
}
510+
raw_keys = self._base_keys(
511+
enable_single_controller=True,
512+
enable_multi_tier_checkpointing=True,
513+
elastic_enabled=True,
514+
mtc_data_parallelism=0,
515+
num_slices=2,
516+
)
517+
518+
max_utils.maybe_initialize_jax_distributed_system(raw_keys)
519+
520+
mock_init.assert_not_called()
521+
mock_mtc_init_kwargs.assert_called_once_with(raw_keys)
522+
mock_mtc.assert_called_once_with(
523+
local_checkpoint_directory=self._base_keys()["local_checkpoint_directory"],
524+
backup_interval_minutes=self._base_keys()["multi_tier_checkpointing_backup_interval_minutes"],
525+
run_name=self._base_keys()["run_name"],
526+
jax_initialization_timeout_seconds=self._base_keys()["jax_distributed_initialization_timeout"],
527+
data_parallelism=1,
528+
num_slices=1,
529+
use_colocated_python=True,
530+
devices=active_devices,
531+
)
532+
495533
@mock.patch("jax.distributed.initialize")
496534
def test_tpu_checkpointing_no_emergency_calls_jax_init(self, mock_init):
497535
raw_keys = self._base_keys(enable_checkpointing=True, compile_topology_num_slices=-1)

0 commit comments

Comments
 (0)