Skip to content

Commit dd150c1

Browse files
authored
[ckpt] fix: Publish checkpoints after tokenizer assets (#5758)
Signed-off-by: Yu Yao <yaoyu.094@gmail.com>
1 parent 35318ee commit dd150c1

3 files changed

Lines changed: 149 additions & 17 deletions

File tree

src/megatron/bridge/training/checkpointing.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1893,6 +1893,8 @@ def save_tokenizer_assets(
18931893
tokenizer: MegatronTokenizer,
18941894
tokenizer_config: TokenizerConfig,
18951895
checkpoint_path: str,
1896+
*,
1897+
raise_on_error: bool = False,
18961898
) -> None:
18971899
"""Save tokenizer files to the checkpoint directory.
18981900
@@ -1904,6 +1906,7 @@ def save_tokenizer_assets(
19041906
tokenizer: The tokenizer instance to save.
19051907
tokenizer_config: The tokenizer configuration (used for file-based tokenizers).
19061908
checkpoint_path: The checkpoint directory path.
1909+
raise_on_error: Propagate tokenizer persistence errors to the caller.
19071910
"""
19081911
if tokenizer is None:
19091912
return
@@ -2022,6 +2025,8 @@ def resolve_path(path_str: str) -> str:
20222025
import traceback
20232026

20242027
logger.error(traceback.format_exc())
2028+
if raise_on_error:
2029+
raise
20252030

20262031

20272032
def _generate_model_state_dict(

src/megatron/bridge/training/model_load_save.py

Lines changed: 26 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -605,6 +605,32 @@ def save_megatron_model(
605605
dist=None,
606606
)
607607

608+
# Complete tokenizer construction and persistence before save_checkpoint publishes
609+
# the root selectors for this checkpoint.
610+
if tokenizer_config is not None:
611+
from megatron.bridge.training.checkpointing import (
612+
get_checkpoint_name,
613+
save_tokenizer_assets,
614+
)
615+
616+
tokenizer_error: Exception | None = None
617+
try:
618+
tokenizer = build_tokenizer(tokenizer_config)
619+
checkpoint_name = get_checkpoint_name(str(path), 0, release=False)
620+
save_tokenizer_assets(tokenizer, tokenizer_config, checkpoint_name, raise_on_error=True)
621+
except Exception as error:
622+
tokenizer_error = error
623+
624+
if torch.distributed.is_initialized():
625+
tokenizer_errors: list[str | None] = [None] * torch.distributed.get_world_size()
626+
local_error = None if tokenizer_error is None else f"{type(tokenizer_error).__name__}: {tokenizer_error}"
627+
torch.distributed.all_gather_object(tokenizer_errors, local_error)
628+
failures = [error for error in tokenizer_errors if error is not None]
629+
if failures:
630+
raise RuntimeError(f"Failed to save tokenizer assets on one or more ranks: {failures}")
631+
elif tokenizer_error is not None:
632+
raise tokenizer_error
633+
608634
if low_memory_save:
609635
# Low-memory save flow: process factories incrementally, freeing memory as we go
610636
import gc
@@ -776,22 +802,6 @@ def _collect_factories(d):
776802
callback_manager=None,
777803
)
778804

779-
# Save tokenizer files separately if tokenizer config is provided
780-
if tokenizer_config is not None:
781-
from megatron.bridge.training.checkpointing import (
782-
get_checkpoint_name,
783-
save_tokenizer_assets,
784-
)
785-
786-
# Build the tokenizer
787-
tokenizer = build_tokenizer(tokenizer_config)
788-
789-
# Get the checkpoint name for step 0
790-
checkpoint_name = get_checkpoint_name(str(path), 0, release=False)
791-
792-
# Save tokenizer files
793-
save_tokenizer_assets(tokenizer, tokenizer_config, checkpoint_name)
794-
795805

796806
def dtype_from_str(dtype: str) -> torch.dtype:
797807
"""Convert a string representation of a dtype to a torch.dtype.

tests/unit_tests/training/test_model_load_save.py

Lines changed: 118 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)