Skip to content

Commit 83527b8

Browse files
Merge pull request AI-Hypercomputer#4407 from AI-Hypercomputer:sujinesh/avoid-mxla-hang-eager-scale-up-upstream
PiperOrigin-RevId: 949675411
2 parents 1942d0e + e89db7a commit 83527b8

3 files changed

Lines changed: 25 additions & 2 deletions

File tree

src/maxtext/common/checkpointing.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1094,8 +1094,8 @@ def maybe_save_checkpoint(checkpoint_manager, state, config, data_iterator, step
10941094
checkpoint_saved = save_checkpoint(checkpoint_manager, actual_step, state, config, data_iterator, force_ckpt_save)
10951095
if checkpoint_saved:
10961096
print_save_message(actual_step, config.async_checkpointing)
1097-
if config.elastic_enabled:
1098-
elastic_utils.maybe_elastic_scale_up(config, checkpoint_manager)
1097+
if config.elastic_enabled:
1098+
elastic_utils.maybe_elastic_scale_up(config, checkpoint_manager)
10991099
except elastic_utils.manager.ScaleUpSignalError as e:
11001100
if config.elastic_enabled:
11011101
max_logging.log(f"Elastic event detected, letting exception bubble up: {e}")

src/maxtext/trainers/pre_train/train.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -680,6 +680,8 @@ def training_loop_iteration(
680680
dump_hlo_upload_all = immutable_data["dump_hlo_upload_all"]
681681

682682
prof.maybe_activate_profiler(step, state)
683+
if config.elastic_enabled:
684+
elastic_utils.maybe_elastic_scale_up(config, checkpoint_manager)
683685

684686
with jax.profiler.StepTraceAnnotation("train", step_num=step):
685687
example_batch = data_loader.load_next_batch(rampup_manager=rampup_manager)

tests/unit/train_state_nnx_checkpoint_test.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -363,6 +363,7 @@ def _config(self, **overrides):
363363
"enable_multi_tier_checkpointing": False,
364364
"local_checkpoint_period": 0,
365365
"enable_autocheckpoint": False,
366+
"elastic_enabled": False,
366367
}
367368
values.update(overrides)
368369
return SimpleNamespace(**values)
@@ -598,6 +599,26 @@ def test_maybe_save_checkpoint_allows_mtc_period_with_continuous_policy(
598599
to_checkpoint_dict_mock.assert_called_once_with(state)
599600
save_checkpoint_mock.assert_called_once()
600601

602+
def test_maybe_save_checkpoint_checks_scale_up_after_unsaved_dispatch(self):
603+
"""Elastic scale-up is checked after save dispatch even when no checkpoint was saved."""
604+
state = mock.Mock()
605+
config = self._config(checkpoint_period=1, elastic_enabled=True)
606+
mgr = mock.MagicMock()
607+
mgr.latest_step.return_value = None
608+
mgr.reached_preemption.return_value = False
609+
save_checkpoint_mock = mock.MagicMock(return_value=False)
610+
611+
with (
612+
mock.patch.object(checkpointing, "save_checkpoint", save_checkpoint_mock),
613+
mock.patch.object(train_state_nnx, "to_checkpoint_dict", return_value={}) as to_checkpoint_dict_mock,
614+
mock.patch.object(checkpointing.elastic_utils, "maybe_elastic_scale_up") as mock_maybe_scale_up,
615+
):
616+
checkpointing.maybe_save_checkpoint(mgr, state, config, data_iterator=None, step=5)
617+
618+
to_checkpoint_dict_mock.assert_called_once_with(state)
619+
save_checkpoint_mock.assert_called_once()
620+
mock_maybe_scale_up.assert_called_once_with(config, mgr)
621+
601622

602623
class TestLinenCheckpointFormatConverters(unittest.TestCase):
603624
"""to_linen_checkpoint_dict / from_linen_checkpoint_dict (NNX <-> Linen on-disk layout)."""

0 commit comments

Comments
 (0)