@@ -1033,9 +1033,126 @@ def finalize(self) -> None:
10331033 mock_build_tokenizer .assert_called_once ()
10341034 mock_get_checkpoint_name .assert_called_once ()
10351035 mock_save_tokenizer_assets .assert_called_once_with (
1036- mock_tokenizer , tokenizer_config , "/fake/checkpoint/iter_0000000"
1036+ mock_tokenizer ,
1037+ tokenizer_config ,
1038+ "/fake/checkpoint/iter_0000000" ,
1039+ raise_on_error = True ,
10371040 )
10381041
1042+ @patch ("megatron.bridge.training.model_load_save.save_checkpoint" )
1043+ @patch ("megatron.bridge.training.model_load_save.get_model_config" )
1044+ @patch ("megatron.bridge.training.model_load_save.GlobalState" )
1045+ @patch ("megatron.bridge.training.model_load_save.ConfigContainer" )
1046+ @patch ("megatron.bridge.training.model_load_save.OptimizerConfig" )
1047+ @patch ("megatron.bridge.training.model_load_save.LoggerConfig" )
1048+ @patch ("megatron.bridge.training.model_load_save.CheckpointConfig" )
1049+ def test_tokenizer_failure_does_not_publish_incomplete_checkpoint (
1050+ self ,
1051+ mock_ckpt_config ,
1052+ mock_logger_config ,
1053+ mock_opt_config ,
1054+ mock_config_container ,
1055+ mock_global_state ,
1056+ mock_get_model_config ,
1057+ mock_save_checkpoint ,
1058+ tmp_path ,
1059+ ):
1060+ """A failed tokenizer save must leave automatic resume on the previous checkpoint."""
1061+
1062+ class MockModelConfig (ModelProviderMixin , Mock ):
1063+ def provide (self , pre_process = None , post_process = None , vp_stage = None ):
1064+ return Mock ()
1065+
1066+ def finalize (self ) -> None :
1067+ pass
1068+
1069+ mock_get_model_config .return_value = MockModelConfig ()
1070+ mock_global_state .return_value = Mock ()
1071+ mock_config_container .return_value = Mock ()
1072+
1073+ latest_train_state = tmp_path / "latest_train_state.pt"
1074+ latest_train_state .write_text ("500" )
1075+
1076+ def publish_selector (** kwargs ):
1077+ latest_train_state .write_text ("0" )
1078+
1079+ mock_save_checkpoint .side_effect = publish_selector
1080+
1081+ tokenizer = Mock ()
1082+ tokenizer .save_pretrained .side_effect = OSError ("tokenizer write failed" )
1083+ checkpoint_name = tmp_path / "iter_0000000"
1084+
1085+ with (
1086+ patch ("megatron.bridge.training.model_load_save.build_tokenizer" , return_value = tokenizer ),
1087+ patch (
1088+ "megatron.bridge.training.checkpointing.get_checkpoint_name" ,
1089+ return_value = str (checkpoint_name ),
1090+ ),
1091+ pytest .raises (OSError , match = "tokenizer write failed" ),
1092+ ):
1093+ save_megatron_model (
1094+ [Mock ()],
1095+ tmp_path ,
1096+ ckpt_format = "torch_dist" ,
1097+ hf_tokenizer_path = "org/model" ,
1098+ low_memory_save = False ,
1099+ )
1100+
1101+ assert latest_train_state .read_text () == "500"
1102+
1103+ @patch ("megatron.bridge.training.model_load_save.save_checkpoint" )
1104+ @patch ("megatron.bridge.training.model_load_save.get_model_config" )
1105+ @patch ("megatron.bridge.training.model_load_save.GlobalState" )
1106+ @patch ("megatron.bridge.training.model_load_save.ConfigContainer" )
1107+ @patch ("megatron.bridge.training.model_load_save.OptimizerConfig" )
1108+ @patch ("megatron.bridge.training.model_load_save.LoggerConfig" )
1109+ @patch ("megatron.bridge.training.model_load_save.CheckpointConfig" )
1110+ def test_tokenizer_failure_stops_all_ranks_before_checkpoint_save (
1111+ self ,
1112+ mock_ckpt_config ,
1113+ mock_logger_config ,
1114+ mock_opt_config ,
1115+ mock_config_container ,
1116+ mock_global_state ,
1117+ mock_get_model_config ,
1118+ mock_save_checkpoint ,
1119+ tmp_path ,
1120+ ):
1121+ """Every rank must observe a tokenizer failure before entering checkpoint save."""
1122+
1123+ class MockModelConfig (ModelProviderMixin , Mock ):
1124+ def provide (self , pre_process = None , post_process = None , vp_stage = None ):
1125+ return Mock ()
1126+
1127+ def finalize (self ) -> None :
1128+ pass
1129+
1130+ mock_get_model_config .return_value = MockModelConfig ()
1131+ mock_global_state .return_value = Mock ()
1132+ mock_config_container .return_value = Mock ()
1133+
1134+ def gather_rank_zero_error (errors , local_error ):
1135+ assert local_error is None
1136+ errors [:] = ["OSError: tokenizer write failed" , None ]
1137+
1138+ with (
1139+ patch ("megatron.bridge.training.model_load_save.build_tokenizer" , return_value = Mock ()),
1140+ patch ("torch.distributed.is_initialized" , return_value = True ),
1141+ patch ("torch.distributed.get_rank" , return_value = 1 ),
1142+ patch ("torch.distributed.get_world_size" , return_value = 2 ),
1143+ patch ("torch.distributed.all_gather_object" , side_effect = gather_rank_zero_error ),
1144+ pytest .raises (RuntimeError , match = "tokenizer write failed" ),
1145+ ):
1146+ save_megatron_model (
1147+ [Mock ()],
1148+ tmp_path ,
1149+ ckpt_format = "torch_dist" ,
1150+ hf_tokenizer_path = "org/model" ,
1151+ low_memory_save = False ,
1152+ )
1153+
1154+ mock_save_checkpoint .assert_not_called ()
1155+
10391156 @patch ("megatron.bridge.training.model_load_save.save_checkpoint" )
10401157 @patch ("megatron.bridge.training.model_load_save.get_model_config" )
10411158 @patch ("megatron.bridge.training.model_load_save.GlobalState" )
0 commit comments