From 26c4e6f2e16ef55d2f888736a85d2efb2148ab88 Mon Sep 17 00:00:00 2001 From: Hejian Sang Date: Wed, 22 Jul 2026 00:18:33 +0000 Subject: [PATCH] fix(megatron): finalize async saves on all ranks before hooks Signed-off-by: Hejian Sang --- miles/backends/megatron_utils/actor.py | 10 ++-- .../megatron_utils/test_actor_checkpoint.py | 48 +++++++++++++++++++ 2 files changed, 54 insertions(+), 4 deletions(-) create mode 100644 tests/fast/backends/megatron_utils/test_actor_checkpoint.py diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 652ff50706..2008f55705 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -613,10 +613,12 @@ def save_model(self, rollout_id: int, force_sync: bool = False) -> None: save_hf_model(self.args, rollout_id, self.model) - if self.args.custom_megatron_post_save_hook_path is not None and dist.get_rank() == 0: - if self.args.async_save: - maybe_finalize_async_save(blocking=True) + post_save_hook_path = self.args.custom_megatron_post_save_hook_path + if post_save_hook_path is not None and self.args.async_save: + # Distributed checkpoint finalization must run on every rank. + maybe_finalize_async_save(blocking=True) + if post_save_hook_path is not None and dist.get_rank() == 0: from megatron.training.checkpointing import get_checkpoint_name from miles.utils.misc import load_function @@ -627,7 +629,7 @@ def save_model(self, rollout_id: int, force_sync: bool = False) -> None: if self.args.save_hf is not None and self.role == "actor" else None ) - post_save_hook = load_function(self.args.custom_megatron_post_save_hook_path) + post_save_hook = load_function(post_save_hook_path) post_save_hook(self.args, rollout_id, checkpoint_dir, hf_checkpoint_dir) if self.args.offload_train: diff --git a/tests/fast/backends/megatron_utils/test_actor_checkpoint.py b/tests/fast/backends/megatron_utils/test_actor_checkpoint.py new file mode 100644 index 0000000000..87ae56c428 --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_actor_checkpoint.py @@ -0,0 +1,48 @@ +from argparse import Namespace +from unittest.mock import MagicMock, patch + +import pytest + +from miles.backends.megatron_utils.actor import MegatronTrainRayActor + + +@pytest.mark.parametrize( + ("rank", "expected_events"), + [ + (0, ["finalize", "save", "finalize", "hook"]), + (1, ["finalize", "save", "finalize"]), + ], +) +def test_async_post_save_hook_finalizes_on_every_rank(rank: int, expected_events: list[str]) -> None: + events: list[str] = [] + actor = MegatronTrainRayActor.__new__(MegatronTrainRayActor) + actor._heartbeat = MagicMock() + actor.args = Namespace( + async_save=True, + custom_megatron_post_save_hook_path="test_checkpoint.post_save_hook", + debug_rollout_only=False, + offload_train=False, + save="/checkpoints", + save_hf=None, + ) + actor.model = MagicMock() + actor.optimizer = MagicMock() + actor.opt_param_scheduler = MagicMock() + actor.role = "actor" + + def post_save_hook(*_args: object) -> None: + events.append("hook") + + with ( + patch("miles.backends.megatron_utils.actor.is_multi_lora_enabled", return_value=False), + patch("miles.backends.megatron_utils.actor.save", side_effect=lambda *_args: events.append("save")), + patch("miles.backends.megatron_utils.actor.dist.get_rank", return_value=rank), + patch( + "megatron.training.async_utils.maybe_finalize_async_save", + side_effect=lambda **_kwargs: events.append("finalize"), + ), + patch("miles.utils.misc.load_function", return_value=post_save_hook), + ): + actor.save_model(rollout_id=7) + + assert events == expected_events