@@ -723,6 +723,65 @@ def test_checkpoint_load_error_propagates(self, mock_ocp):
723723 with self .assertRaises (RuntimeError ):
724724 model_creation_utils .from_pretrained (cfg , self .mesh )
725725
726+ @patch ("maxtext.utils.model_creation_utils.checkpointing.load_checkpoint_metadata" )
727+ def test_scan_layers_mismatch_raises_error (self , mock_load_meta ):
728+ """ValueError is raised if run specifies scan_layers=True but checkpoint specifies scan_layers=False."""
729+ mock_load_meta .return_value = {"scan_layers" : False }
730+
731+ cfg = _make_config (
732+ enable_checkpointing = True , load_parameters_path = "gs://fake/scan_layers_false_ckpt" , scan_layers = True
733+ )
734+
735+ with self .assertRaises (ValueError ) as context :
736+ model_creation_utils .from_pretrained (cfg , self .mesh )
737+ self .assertIn (
738+ "Configuration mismatch: Your run specifies scan_layers=True, "
739+ "but the checkpoint was saved with scan_layers=False" ,
740+ str (context .exception ),
741+ )
742+
743+ @patch ("maxtext.utils.model_creation_utils.checkpointing.load_checkpoint_metadata" )
744+ @patch ("maxtext.utils.model_creation_utils.ocp" )
745+ def test_scan_layers_match_no_error (self , mock_ocp , mock_load_meta ):
746+ """If the run specifies scan_layers=True and the checkpoint matches, it proceeds without error."""
747+ mock_load_meta .return_value = {"scan_layers" : True }
748+
749+ mock_ckptr = MagicMock ()
750+ mock_ckptr .metadata .return_value = self ._make_linen_metadata_mock ()
751+ mock_ckptr .restore .side_effect = lambda path , item = None , ** kw : item
752+ mock_ocp .Checkpointer .return_value = mock_ckptr
753+ mock_ocp .PyTreeCheckpointHandler .return_value = MagicMock ()
754+ mock_ocp .checkpoint_utils .construct_restore_args .return_value = {}
755+ mock_ocp .ArrayRestoreArgs = ocp .ArrayRestoreArgs
756+
757+ cfg = _make_config (
758+ enable_checkpointing = True , load_parameters_path = "gs://fake/scan_layers_true_ckpt" , scan_layers = True
759+ )
760+
761+ model = model_creation_utils .from_pretrained (cfg , self .mesh )
762+ self .assertIsInstance (model , models .Transformer )
763+
764+ @patch ("maxtext.utils.model_creation_utils.checkpointing.load_checkpoint_metadata" )
765+ @patch ("maxtext.utils.model_creation_utils.ocp" )
766+ def test_scan_layers_missing_metadata_no_error (self , mock_ocp , mock_load_meta ):
767+ """Skip verification and proceed if custom_metadata lacks 'scan_layers'."""
768+ mock_load_meta .return_value = {}
769+
770+ mock_ckptr = MagicMock ()
771+ mock_ckptr .metadata .return_value = self ._make_linen_metadata_mock ()
772+ mock_ckptr .restore .side_effect = lambda path , item = None , ** kw : item
773+ mock_ocp .Checkpointer .return_value = mock_ckptr
774+ mock_ocp .PyTreeCheckpointHandler .return_value = MagicMock ()
775+ mock_ocp .checkpoint_utils .construct_restore_args .return_value = {}
776+ mock_ocp .ArrayRestoreArgs = ocp .ArrayRestoreArgs
777+
778+ cfg = _make_config (
779+ enable_checkpointing = True , load_parameters_path = "gs://fake/scan_layers_missing_ckpt" , scan_layers = True
780+ )
781+
782+ model = model_creation_utils .from_pretrained (cfg , self .mesh )
783+ self .assertIsInstance (model , models .Transformer )
784+
726785
727786class TestSetupDecodeStateFromNnx (unittest .TestCase ):
728787 """Tests for setup_decode_state_from_nnx()."""
0 commit comments