diff --git a/src/megatron/bridge/training/optim.py b/src/megatron/bridge/training/optim.py index 6d7df9d45f..4947d203e5 100644 --- a/src/megatron/bridge/training/optim.py +++ b/src/megatron/bridge/training/optim.py @@ -13,8 +13,12 @@ # limitations under the License. import logging -from typing import Optional, Union +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager +from types import MethodType +from typing import Optional, Union, cast +import torch from megatron.core.optimizer import ( MegatronOptimizer, OptimizerConfig, @@ -111,6 +115,132 @@ def setup_optimizer( return optimizer, scheduler +def _optimizer_state_is_fp32_adam(state_dict: Mapping[str, object]) -> bool: + """Return whether a state dict contains only FP32 Adam moment tensors.""" + states = state_dict.get("state") + if not isinstance(states, Mapping) or not states: + return False + + expected_state_names = {"exp_avg", "exp_avg_sq"} + for state in states.values(): + if not isinstance(state, Mapping) or set(state) != expected_state_names: + return False + if any(not isinstance(value, torch.Tensor) or value.dtype != torch.float32 for value in state.values()): + return False + return True + + +def _get_te_fused_adam_class() -> type[torch.optim.Optimizer] | None: + """Return Transformer Engine's FusedAdam class when it is available.""" + try: + from transformer_engine.pytorch.optimizers import FusedAdam + except ImportError: + return None + return cast(type[torch.optim.Optimizer], FusedAdam) + + +@contextmanager +def memory_efficient_fp32_optimizer_state_loading( + optimizer: MegatronOptimizer | None, +) -> Iterator[int]: + """Avoid redundant TE FusedAdam state copies during checkpoint loading. + + Megatron-Core's distributed optimizer creates FP32 Adam moment tensors as + its loading scaffold. Transformer Engine's loader replaces those tensors + with another FP32 copy even though no conversion is needed. This context + temporarily uses PyTorch's base loader only for compatible standard + distributed FusedAdam instances, allowing them to adopt the scaffold + tensors directly. + + Args: + optimizer: Optimizer participating in checkpoint loading. + + Yields: + Number of compatible FusedAdam instances using the base loader. + """ + if optimizer is None: + yield 0 + return + + fused_adam_class = _get_te_fused_adam_class() + if fused_adam_class is None: + yield 0 + return + + sub_optimizers = optimizer.chained_optimizers if hasattr(optimizer, "chained_optimizers") else [optimizer] + missing_method = object() + patched: list[tuple[torch.optim.Optimizer, object]] = [] + + try: + for distributed_optimizer in sub_optimizers: + if getattr(distributed_optimizer, "is_stub_optimizer", False): + continue + if not hasattr(distributed_optimizer, "shard_fp32_from_float16_groups"): + continue + if getattr(getattr(distributed_optimizer, "ddp_config", None), "use_megatron_fsdp", False): + continue + + config = getattr(distributed_optimizer, "config", None) + if getattr(config, "use_precision_aware_optimizer", False): + continue + if getattr(config, "optimizer_cpu_offload", False): + continue + + inner = getattr(distributed_optimizer, "optimizer", None) + if not isinstance(inner, fused_adam_class): + continue + if getattr(inner, "master_weights", None) is not False: + continue + if getattr(inner, "store_param_remainders", False): + continue + + state_dtype_map = getattr(inner, "name_to_dtype_map", None) + if not isinstance(state_dtype_map, Mapping) or any( + state_dtype_map.get(name) != torch.float32 for name in ("exp_avg", "exp_avg_sq") + ): + continue + + params = [param for group in inner.param_groups for param in group["params"]] + if not params or any(param.dtype != torch.float32 for param in params): + continue + + original_load_state_dict: Callable[[dict[str, object]], None] = inner.load_state_dict + + def _load_state_dict_without_fp32_reallocation( + fused_adam: torch.optim.Optimizer, + state_dict: dict[str, object], + *, + _fallback: Callable[[dict[str, object]], None] = original_load_state_dict, + ) -> None: + if not _optimizer_state_is_fp32_adam(state_dict): + _fallback(state_dict) + return + torch.optim.Optimizer.load_state_dict(fused_adam, state_dict) + + previous_instance_method = inner.__dict__.get("load_state_dict", missing_method) + setattr(inner, "load_state_dict", MethodType(_load_state_dict_without_fp32_reallocation, inner)) + patched.append((inner, previous_instance_method)) + + if patched: + G_LOGGER.info( + "Enabled memory-efficient FP32 checkpoint-state loading for %d distributed " + "Transformer Engine FusedAdam optimizer(s).", + len(patched), + ) + + yield len(patched) + finally: + for inner, previous_instance_method in patched: + if previous_instance_method is missing_method: + delattr(inner, "load_state_dict") + else: + setattr(inner, "load_state_dict", previous_instance_method) + if patched: + # The replaced scaffold tensors cannot enter the allocator cache + # until the checkpoint-loading frame releases its state dict. + torch.cuda.empty_cache() + + def sync_hybrid_device_optimizer_fp32_master_copies(optimizer: MegatronOptimizer | None) -> bool: """Refresh ``HybridDeviceOptimizer`` FP32 master copies from BF16 model parameters. diff --git a/src/megatron/bridge/training/setup.py b/src/megatron/bridge/training/setup.py index 3f1506e8a2..bbe03ad2f8 100644 --- a/src/megatron/bridge/training/setup.py +++ b/src/megatron/bridge/training/setup.py @@ -46,7 +46,11 @@ ) from megatron.bridge.training.config import ConfigContainer from megatron.bridge.training.initialize import initialize_megatron, set_jit_fusion_options -from megatron.bridge.training.optim import setup_optimizer, sync_hybrid_device_optimizer_fp32_master_copies +from megatron.bridge.training.optim import ( + memory_efficient_fp32_optimizer_state_loading, + setup_optimizer, + sync_hybrid_device_optimizer_fp32_master_copies, +) from megatron.bridge.training.state import GlobalState from megatron.bridge.training.tensor_inspect import ( finalize_tensor_inspect_post_model_initialization, @@ -346,15 +350,17 @@ def modelopt_pre_wrap_hook(model): if should_load_checkpoint: timers("load-checkpoint", log_level=0).start(barrier=True) - checkpoint_manager.load( - CheckpointLoadContext( - state=state, - model=model, - optimizer=optimizer, - opt_param_scheduler=scheduler, - skip_load_to_model_and_opt=cfg.dist.use_torch_fsdp2, + checkpoint_optimizer = optimizer if cfg.checkpoint.load_optim and not cfg.checkpoint.finetune else None + with memory_efficient_fp32_optimizer_state_loading(checkpoint_optimizer): + checkpoint_manager.load( + CheckpointLoadContext( + state=state, + model=model, + optimizer=optimizer, + opt_param_scheduler=scheduler, + skip_load_to_model_and_opt=cfg.dist.use_torch_fsdp2, + ) ) - ) # Workaround for upstream mcore: reload_model_params() only refreshes the # level-1 FP32 GPU shards of HybridDeviceOptimizer, so the level-2 CPU # clones and level-3 FP32 working copies retain their random init. Without diff --git a/tests/unit_tests/training/test_optim.py b/tests/unit_tests/training/test_optim.py index 8ea17a27c6..e00b02abb5 100644 --- a/tests/unit_tests/training/test_optim.py +++ b/tests/unit_tests/training/test_optim.py @@ -15,13 +15,18 @@ """Tests for setup_optimizer in optim.py.""" import builtins +from types import SimpleNamespace from unittest.mock import MagicMock, patch +import pytest import torch from megatron.core.optimizer import OptimizerConfig, ParamGroupOverride, ParamKey from megatron.bridge.training.config import SchedulerConfig -from megatron.bridge.training.optim import sync_hybrid_device_optimizer_fp32_master_copies +from megatron.bridge.training.optim import ( + memory_efficient_fp32_optimizer_state_loading, + sync_hybrid_device_optimizer_fp32_master_copies, +) class TestSetupOptimizerMuP: @@ -155,6 +160,27 @@ class _FakeHDO: """Stand-in for HybridDeviceOptimizer used to satisfy the isinstance check.""" +class _FakeFusedAdam(torch.optim.Optimizer): + """CPU stand-in for TE FusedAdam with an observable override loader.""" + + def __init__( + self, + param: torch.Tensor, + *, + master_weights: bool = False, + exp_avg_dtype: torch.dtype = torch.float32, + ) -> None: + super().__init__([param], {"lr": 1e-3}) + self.master_weights = master_weights + self.store_param_remainders = False + self.name_to_dtype_map = {"exp_avg": exp_avg_dtype, "exp_avg_sq": exp_avg_dtype} + self.override_load_calls = 0 + + def load_state_dict(self, state_dict: dict[str, object]) -> None: + self.override_load_calls += 1 + super().load_state_dict(state_dict) + + class _FakeParamRange: def __init__(self, start: int, end: int): self.start = start @@ -169,6 +195,9 @@ def __init__(self, *, model_param: torch.Tensor, shard_main_param: torch.Tensor self.model_float16_groups = [[model_param]] self.shard_fp32_from_float16_groups = [[shard_main_param]] self._numel = model_param.numel() + self.is_stub_optimizer = False + self.ddp_config = SimpleNamespace(use_megatron_fsdp=False) + self.config = SimpleNamespace(use_precision_aware_optimizer=False, optimizer_cpu_offload=False) def _get_model_param_range_map(self, _param: torch.Tensor) -> dict: return {"param": _FakeParamRange(0, self._numel)} @@ -188,6 +217,211 @@ def __init__(self, sub_opts: list[object]) -> None: self.chained_optimizers = sub_opts +class _FakeLayerWiseChildOpt: + """Stand-in for a LayerWiseDistributedOptimizer's wrapped child optimizer.""" + + def __init__(self, inner: torch.optim.Optimizer) -> None: + self.optimizer = inner + + +class TestMemoryEfficientFp32OptimizerStateLoading: + """Tests for the scoped TE FusedAdam checkpoint-load fast path.""" + + @staticmethod + def _distributed_optimizer( + *, + param_dtype: torch.dtype = torch.float32, + master_weights: bool = False, + state_dtype: torch.dtype = torch.float32, + ) -> tuple[_FakeDistribOpt, _FakeFusedAdam, torch.Tensor]: + param = torch.zeros(4, dtype=param_dtype) + inner = _FakeFusedAdam(param, master_weights=master_weights, exp_avg_dtype=state_dtype) + distributed = _FakeDistribOpt( + model_param=torch.zeros(4, dtype=torch.bfloat16), + shard_main_param=param, + inner=inner, + ) + return distributed, inner, param + + @staticmethod + def _state_dict( + inner: _FakeFusedAdam, + *, + dtype: torch.dtype = torch.float32, + ) -> tuple[dict[str, object], torch.Tensor]: + state_dict = inner.state_dict() + exp_avg = torch.ones(4, dtype=dtype) + state_dict["state"] = { + 0: { + "exp_avg": exp_avg, + "exp_avg_sq": torch.full((4,), 2.0, dtype=dtype), + } + } + return state_dict, exp_avg + + def test_uses_base_loader_without_reallocating_fp32_state(self): + """FP32 distributed shards adopt supplied state tensors directly.""" + distributed, inner, param = self._distributed_optimizer() + state_dict, exp_avg = self._state_dict(inner) + + with ( + patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam), + patch("megatron.bridge.training.optim.torch.cuda.empty_cache") as mock_empty_cache, + ): + with memory_efficient_fp32_optimizer_state_loading(distributed) as patched: + inner.load_state_dict(state_dict) + mock_empty_cache.assert_not_called() + + assert patched == 1 + assert inner.override_load_calls == 0 + assert inner.state[param]["exp_avg"] is exp_avg + mock_empty_cache.assert_called_once_with() + + inner.load_state_dict(state_dict) + + assert inner.override_load_calls == 1 + + def test_falls_back_for_non_fp32_incoming_state(self): + """A non-FP32 state dict retains Transformer Engine's conversion path.""" + distributed, inner, _ = self._distributed_optimizer() + state_dict, _ = self._state_dict(inner, dtype=torch.bfloat16) + + with ( + patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam), + patch("megatron.bridge.training.optim.torch.cuda.empty_cache"), + memory_efficient_fp32_optimizer_state_loading(distributed) as patched, + ): + inner.load_state_dict(state_dict) + + assert patched == 1 + assert inner.override_load_calls == 1 + + @pytest.mark.parametrize( + ("param_dtype", "master_weights", "state_dtype"), + [ + (torch.bfloat16, False, torch.float32), + (torch.float32, True, torch.float32), + (torch.float32, False, torch.bfloat16), + ], + ) + def test_does_not_patch_incompatible_fused_adam( + self, + param_dtype: torch.dtype, + master_weights: bool, + state_dtype: torch.dtype, + ): + """Mixed parameters, master weights, and compressed state stay on TE's path.""" + distributed, inner, _ = self._distributed_optimizer( + param_dtype=param_dtype, + master_weights=master_weights, + state_dtype=state_dtype, + ) + + with patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam): + with memory_efficient_fp32_optimizer_state_loading(distributed) as patched: + pass + + assert patched == 0 + assert "load_state_dict" not in inner.__dict__ + + @pytest.mark.parametrize("incompatibility", ["precision_aware", "cpu_offload", "fsdp", "stub"]) + def test_does_not_patch_incompatible_distributed_optimizer(self, incompatibility: str): + """Special distributed optimizer modes retain their existing loader.""" + distributed, inner, _ = self._distributed_optimizer() + if incompatibility == "precision_aware": + distributed.config.use_precision_aware_optimizer = True + elif incompatibility == "cpu_offload": + distributed.config.optimizer_cpu_offload = True + elif incompatibility == "fsdp": + distributed.ddp_config.use_megatron_fsdp = True + else: + distributed.is_stub_optimizer = True + + with patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam): + with memory_efficient_fp32_optimizer_state_loading(distributed) as patched: + pass + + assert patched == 0 + assert "load_state_dict" not in inner.__dict__ + + def test_patches_all_eligible_chained_optimizers(self): + """Dense and expert DistributedOptimizers both use the scoped loader.""" + distributed_optimizers = [self._distributed_optimizer()[0] for _ in range(2)] + + with ( + patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam), + patch("megatron.bridge.training.optim.torch.cuda.empty_cache"), + ): + with memory_efficient_fp32_optimizer_state_loading(_ChainedOpt(distributed_optimizers)) as patched: + assert all("load_state_dict" in opt.optimizer.__dict__ for opt in distributed_optimizers) + + assert patched == 2 + assert all("load_state_dict" not in opt.optimizer.__dict__ for opt in distributed_optimizers) + + def test_restores_methods_when_later_optimizer_setup_raises(self): + """A partial chained-optimizer setup is rolled back when inspection fails.""" + first_distributed, first_inner, _ = self._distributed_optimizer() + second_distributed, second_inner, _ = self._distributed_optimizer() + second_inner.param_groups = [{}] + + with ( + patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam), + patch("megatron.bridge.training.optim.torch.cuda.empty_cache") as mock_empty_cache, + pytest.raises(KeyError, match="params"), + ): + with memory_efficient_fp32_optimizer_state_loading(_ChainedOpt([first_distributed, second_distributed])): + pass + + assert "load_state_dict" not in first_inner.__dict__ + mock_empty_cache.assert_called_once_with() + + def test_te_unavailable_is_noop(self): + """An environment without Transformer Engine retains the existing loader.""" + distributed, inner, _ = self._distributed_optimizer() + + with ( + patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=None), + patch("megatron.bridge.training.optim.torch.cuda.empty_cache") as mock_empty_cache, + memory_efficient_fp32_optimizer_state_loading(distributed) as patched, + ): + pass + + assert patched == 0 + assert "load_state_dict" not in inner.__dict__ + mock_empty_cache.assert_not_called() + + def test_does_not_patch_layerwise_optimizer_children(self): + """LayerWise optimizer children lack distributed FP32 shards and remain unchanged.""" + _, inner, _ = self._distributed_optimizer() + layerwise = _ChainedOpt([_FakeLayerWiseChildOpt(inner)]) + + with patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam): + with memory_efficient_fp32_optimizer_state_loading(layerwise) as patched: + pass + + assert patched == 0 + assert "load_state_dict" not in inner.__dict__ + + def test_restores_methods_when_loading_raises(self): + """The scoped replacement is removed when checkpoint loading fails.""" + distributed, inner, _ = self._distributed_optimizer() + + with ( + patch("megatron.bridge.training.optim._get_te_fused_adam_class", return_value=_FakeFusedAdam), + patch("megatron.bridge.training.optim.torch.cuda.empty_cache"), + pytest.raises(RuntimeError, match="load failed"), + ): + with memory_efficient_fp32_optimizer_state_loading(distributed): + raise RuntimeError("load failed") + + assert "load_state_dict" not in inner.__dict__ + + def test_none_optimizer_is_noop(self): + """A missing optimizer is a no-op.""" + with memory_efficient_fp32_optimizer_state_loading(None) as patched: + assert patched == 0 + + class TestSyncHybridDeviceOptimizerFp32MasterCopies: """Tests for the post-load FP32 master sync workaround helper."""