diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 8d47d6a1793..78d7e41e91f 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -893,14 +893,3 @@ def sharded_state_dict( nonempty_rank_group['params'] = local_params sd['optimizer']['param_groups'][i] = nonempty_rank_group return sharded_state_dict - - def save_state_dict_to_file(self, filename: str) -> None: - """Save the parameter state of the optimizer. For torch format only. - Args: - filename: The filename to save the parameter state. - """ - torch.save(super().state_dict(), filename) - - def load_state_dict_from_file(self, filename: str) -> None: - """Load the parameter state of the optimizer. For torch format only.""" - super().load_state_dict(torch.load(filename)) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index f11d98e309d..52166c0cc18 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -1375,6 +1375,14 @@ def state_dict(self): else: return [optimizer.state_dict() for optimizer in self.chained_optimizers] + def save_state_dict_to_file(self, filename: str) -> None: + """Save this optimizer's per-rank state for torch checkpoints.""" + torch.save(self.state_dict(), filename) + + def load_state_dict_from_file(self, filename: str) -> None: + """Load this optimizer's per-rank state from a torch checkpoint.""" + self.load_state_dict(torch.load(filename)) + def sharded_state_dict( self, model_sharded_state_dict: ShardedStateDict, is_loading: bool = False, **kwargs ): diff --git a/tests/unit_tests/test_optimizer.py b/tests/unit_tests/test_optimizer.py index 3c445cd3633..5b3e69c23b8 100644 --- a/tests/unit_tests/test_optimizer.py +++ b/tests/unit_tests/test_optimizer.py @@ -326,6 +326,46 @@ def to_cuda(d): assert list(optimizer_2.state.values())[0]["momentum_buffer"].is_cuda +def test_chained_optimizer_file_state_dict_round_trip(tmp_path): + """Torch checkpoint files preserve state from every chained optimizer.""" + + class MockOptimizer: + def __init__(self, state_dict): + self.config = None + self.model_chunks = [] + self.is_stub_optimizer = False + self.optimizer = self + self.param_groups = [{'params': []}] + self._state_dict = state_dict + + def state_dict(self): + return self._state_dict + + def load_state_dict(self, state_dict): + self._state_dict = state_dict + + state_dicts = [ + {'optimizer': {'state': {'layer_wise': 'state'}, 'param_groups': []}}, + {'optimizer': {'state': {'distributed': 'state'}, 'param_groups': []}}, + ] + checkpoint_path = tmp_path / f'optimizer_{os.getpid()}.pt' + + ChainedOptimizer( + [MockOptimizer(state_dict) for state_dict in state_dicts] + ).save_state_dict_to_file(checkpoint_path) + + restored = ChainedOptimizer( + [ + MockOptimizer({'optimizer': {'state': {}, 'param_groups': []}}), + MockOptimizer({'optimizer': {'state': {}, 'param_groups': []}}), + ] + ) + restored.load_state_dict_from_file(checkpoint_path) + + assert torch.load(checkpoint_path) == state_dicts + assert restored.state_dict() == state_dicts + + def test_chained_optimizer_get_parameters(): """Test ChainedOptimizer.get_parameters() aggregates params from all sub-optimizers.