@@ -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
602623class TestLinenCheckpointFormatConverters (unittest .TestCase ):
603624 """to_linen_checkpoint_dict / from_linen_checkpoint_dict (NNX <-> Linen on-disk layout)."""
0 commit comments