From 3c5b4bb48979128e8031b48bb595919466e0ca44 Mon Sep 17 00:00:00 2001 From: Chen Cui Date: Tue, 4 Aug 2026 17:24:37 -0700 Subject: [PATCH] fix(checkpointing): honor disabled RNG saves Signed-off-by: Chen Cui --- src/megatron/bridge/training/checkpointing.py | 14 ++--- .../bridge/training/model_load_save.py | 8 +-- .../unit_tests/training/test_checkpointing.py | 7 +++ .../training/test_model_load_save.py | 53 ++++++++++++++++--- 4 files changed, 62 insertions(+), 20 deletions(-) diff --git a/src/megatron/bridge/training/checkpointing.py b/src/megatron/bridge/training/checkpointing.py index ac4508884f..72a594b17e 100644 --- a/src/megatron/bridge/training/checkpointing.py +++ b/src/megatron/bridge/training/checkpointing.py @@ -1152,12 +1152,14 @@ def save_checkpoint( # Collect rng state across data parallel ranks. if pg_collection is None: pg_collection = get_pg_collection(model) - rng_state = get_rng_state( - data_parallel_random_init=cfg.rng.data_parallel_random_init, - ckpt_format=ckpt_cfg.ckpt_format, - pg_collection=pg_collection, - module_name=module_name, - ) + rng_state = None + if ckpt_cfg.save_rng: + rng_state = get_rng_state( + data_parallel_random_init=cfg.rng.data_parallel_random_init, + ckpt_format=ckpt_cfg.ckpt_format, + pg_collection=pg_collection, + module_name=module_name, + ) # Collect rerun state across all ranks rerun_state_machine = get_rerun_state_machine() diff --git a/src/megatron/bridge/training/model_load_save.py b/src/megatron/bridge/training/model_load_save.py index a077b8a520..d9c638b935 100644 --- a/src/megatron/bridge/training/model_load_save.py +++ b/src/megatron/bridge/training/model_load_save.py @@ -622,17 +622,13 @@ def save_megatron_model( from megatron.bridge.training.checkpointing import ( _build_sharded_state_dict_metadata, generate_state_dict, - get_rng_state, ) from megatron.bridge.training.utils.pg_utils import get_pg_collection logger.info("[LOW_MEMORY_SAVE] Generating state dict...") - # Get RNG state (minimal, since save_rng=False) + # Conversion checkpoints intentionally omit RNG state. pg_collection = get_pg_collection(model) - rng_state = get_rng_state( - data_parallel_random_init=False, ckpt_format=ckpt_format, pg_collection=pg_collection - ) # Build sharded state dict metadata sharded_sd_metadata = _build_sharded_state_dict_metadata(False, state.cfg.checkpoint) @@ -643,7 +639,7 @@ def save_megatron_model( model, optimizer=None, opt_param_scheduler=None, - rng_state=rng_state, + rng_state=None, iteration=0, optim_sd_kwargs=dict(metadata=sharded_sd_metadata), model_sd_kwargs=dict(metadata=sharded_sd_metadata), diff --git a/tests/unit_tests/training/test_checkpointing.py b/tests/unit_tests/training/test_checkpointing.py index a47083692e..ac18f0eb98 100644 --- a/tests/unit_tests/training/test_checkpointing.py +++ b/tests/unit_tests/training/test_checkpointing.py @@ -600,6 +600,7 @@ def save_checkpoint_fixtures(): class TestSaveCheckpoint: """Test checkpoint saving functionality.""" + @pytest.mark.parametrize("save_rng", [True, False]) @patch("megatron.bridge.training.checkpointing.wandb_utils") @patch("megatron.bridge.training.checkpointing.is_last_rank") @patch("builtins.open", new_callable=mock_open) @@ -647,6 +648,7 @@ def test_save_checkpoint_global( mock_is_last_rank, mock_wandb, save_checkpoint_fixtures, + save_rng, ): """Test saving a global checkpoint.""" # Setup mocks @@ -678,6 +680,7 @@ def test_save_checkpoint_global( # Add wandb logger to state save_checkpoint_fixtures["mock_state"].wandb_logger = Mock() save_checkpoint_fixtures["mock_state"].cfg.checkpoint.most_recent_k = -1 + save_checkpoint_fixtures["mock_state"].cfg.checkpoint.save_rng = save_rng # Call save_checkpoint save_checkpoint( @@ -694,6 +697,10 @@ def test_save_checkpoint_global( mock_ft.on_checkpointing_start.assert_called_once() mock_gen_state.assert_called_once() mock_dist_ckpt.save.assert_called_once() + if save_rng: + mock_get_rng.assert_called_once() + else: + mock_get_rng.assert_not_called() # Verify that the tracker file was written with the correct iteration tracker_calls = [ diff --git a/tests/unit_tests/training/test_model_load_save.py b/tests/unit_tests/training/test_model_load_save.py index ee93303cb5..0ddb79fe56 100644 --- a/tests/unit_tests/training/test_model_load_save.py +++ b/tests/unit_tests/training/test_model_load_save.py @@ -846,16 +846,53 @@ def test_load_megatron_model_applies_overrides(self, mock_load_model_config, moc class TestSaveMegatronModel: """Test save_megatron_model function. - Note: These tests use low_memory_save=False because the low_memory_save=True path - requires parallel state to be initialized (get_rng_state calls mpu.get_pipeline_model_parallel_rank()). - Testing the low_memory_save=True path would require either: - 1. Full distributed initialization, or - 2. Extensive mocking of checkpointing internals (get_rng_state, generate_state_dict, etc.) - - The low_memory_save=False path tests the core save_checkpoint integration without - those dependencies, which is sufficient for unit testing the function's API and behavior. + Most tests use low_memory_save=False to exercise save_checkpoint integration + without mocking the incremental state-dict processing machinery. """ + def test_low_memory_save_omits_rng_collection(self): + """Low-memory conversion saves must not initialize CUDA for disabled RNG state.""" + + class MockModelConfig(ModelProviderMixin, Mock): + def provide(self, pre_process=None, post_process=None, vp_stage=None): + return Mock() + + def finalize(self) -> None: + pass + + mock_model = Mock() + mock_model.named_parameters.return_value = [] + mock_model.parameters.return_value = [] + mock_pg_collection = Mock() + + with ( + tempfile.TemporaryDirectory() as temp_dir, + patch( + "megatron.bridge.training.model_load_save.get_model_config", + return_value=MockModelConfig(), + ), + patch( + "megatron.bridge.training.utils.pg_utils.get_pg_collection", + return_value=mock_pg_collection, + ), + patch( + "megatron.bridge.training.checkpointing.get_rng_state", + ) as mock_get_rng_state, + patch( + "megatron.bridge.training.checkpointing._build_sharded_state_dict_metadata", + return_value={}, + ), + patch( + "megatron.bridge.training.checkpointing.generate_state_dict", + return_value={}, + ) as mock_generate_state_dict, + patch("megatron.bridge.training.model_load_save.save_checkpoint"), + ): + save_megatron_model([mock_model], temp_dir, ckpt_format="torch_dist", low_memory_save=True) + + mock_get_rng_state.assert_not_called() + assert mock_generate_state_dict.call_args.kwargs["rng_state"] is None + @patch("megatron.bridge.training.model_load_save.save_checkpoint") @patch("megatron.bridge.training.model_load_save.get_model_config") @patch("megatron.bridge.training.model_load_save.GlobalState")