diff --git a/miles/backends/megatron_utils/lora_utils.py b/miles/backends/megatron_utils/lora_utils.py index 7fe3804e7a7..fe875e67aa0 100644 --- a/miles/backends/megatron_utils/lora_utils.py +++ b/miles/backends/megatron_utils/lora_utils.py @@ -11,6 +11,8 @@ import torch.distributed as dist from miles.backends.training_utils.parallel import get_parallel_state +from miles.utils.distributed_utils import get_gloo_group +from miles.utils.ft_utils.process_group_utils import collective_bool_and from miles.utils.lora import is_lora_enabled, lora_rollout_enabled # noqa: F401 (re-exported) logger = logging.getLogger(__name__) @@ -404,6 +406,90 @@ def create_lora_instance(args: Namespace): # --------------------------------------------------------------------------- +def _all_ranks_true(value: bool) -> bool: + if not dist.is_initialized(): + return value + return collective_bool_and(value=value, group=get_gloo_group()) + + +def _raise_if_any_rank_failed(local_error: Exception | None, message: str) -> None: + """Raise on every rank when any rank reports ``local_error``, so failures stay collective.""" + if _all_ranks_true(local_error is None): + return + if local_error is not None: + raise RuntimeError(message) from local_error + raise RuntimeError(message) + + +def _optimizer_param_state_entries(optimizer: Any, directory: Path) -> list[tuple[Any, Path]]: + """``(child, parameter-state file)`` for the children whose state is not in ``state_dict()``. + + ``DistributedOptimizer`` shards master weights and Adam moments over the data-parallel + group and keeps them out of ``state_dict()``; they are saved to their own per-rank files. + Stubs hold no shard at all and are skipped, as they are in Megatron's own load path. + """ + from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer + + rank = dist.get_rank() if dist.is_initialized() else 0 + children = getattr(optimizer, "chained_optimizers", [optimizer]) + return [ + (child, directory / f"optimizer_param_state_rank{rank}_optimizer{index}.pt") + for index, child in enumerate(children) + if isinstance(child, DistributedOptimizer) and not child.is_stub_optimizer + ] + + +def _holds_gathered_param_state(child: Any) -> bool: + """Only the data-parallel root gathers the full state, and so only it reads and writes it.""" + return child.data_parallel_group.rank() == 0 + + +def _save_optimizer_param_state(optimizer: Any, directory: Path) -> None: + """``save_parameter_state`` gathers over the data-parallel group: every rank must call it.""" + save_error = None + for child, path in _optimizer_param_state_entries(optimizer, directory): + try: + child.save_parameter_state(str(path)) + except Exception as error: + save_error = save_error or error + _raise_if_any_rank_failed(save_error, "Failed to save optimizer parameter state on at least one rank") + + +def _check_param_state_matches_buffers(child: Any, state: dict) -> None: + """``load_parameter_state_from_dp_zero`` asserts this between two scatters, where a failure + on the data-parallel root hangs the rest of the group; check it before any collective runs.""" + child.split_state_dict_if_needed(state) + for gbuf_index, dtypes in enumerate(child.gbuf_ranges): + for dtype in dtypes: + expected = child.buffers[gbuf_index].numel_unpadded + found = state[gbuf_index][dtype]["numel_unpadded"] + if expected != found: + raise RuntimeError( + f"Optimizer parameter state does not match the model: buffer {gbuf_index} holds " + f"{expected} unpadded elements, the checkpoint holds {found}" + ) + + +def _load_optimizer_param_state(entries: list[tuple[Any, Path]]) -> None: + """``load_parameter_state_from_dp_zero`` scatters from the data-parallel root: every rank + must call it, so the reads are coordinated before the first scatter.""" + states = [] + load_error = None + for child, path in entries: + state = None + if _holds_gathered_param_state(child): + try: + state = torch.load(path, map_location="cpu", weights_only=True) + _check_param_state_matches_buffers(child, state) + except Exception as error: + load_error = load_error or error + states.append((child, state)) + + _raise_if_any_rank_failed(load_error, "Failed to read optimizer parameter state on at least one rank") + for child, state in states: + child.load_parameter_state_from_dp_zero(state) + + def save_lora_checkpoint( model: Sequence[torch.nn.Module], args: Namespace, @@ -428,7 +514,8 @@ def save_lora_checkpoint( never change, so they are not saved. This function is collective: **all ranks must call it** because the bridge - export performs TP all-gather internally. Only ``dp_rank == 0`` writes files. + export performs TP all-gather internally and the optimizer parameter-state + save gathers over the data-parallel group. """ import json @@ -503,14 +590,21 @@ def save_lora_checkpoint( # ---- Training state (optimizer + scheduler) for resume ---- if optimizer is not None: rank = dist.get_rank() if dist.is_initialized() else 0 - torch.save( - { - "iteration": iteration, - "optimizer": optimizer.state_dict(), - "opt_param_scheduler": opt_param_scheduler.state_dict() if opt_param_scheduler else None, - }, - save_path / f"training_state_rank{rank}.pt", - ) + save_error = None + try: + torch.save( + { + "iteration": iteration, + "optimizer": optimizer.state_dict(), + "opt_param_scheduler": opt_param_scheduler.state_dict() if opt_param_scheduler else None, + }, + save_path / f"training_state_rank{rank}.pt", + ) + except Exception as error: + save_error = error + # The parameter-state save below is collective; a rank that raised here would hang its peers. + _raise_if_any_rank_failed(save_error, "Failed to save optimizer training state on at least one rank") + _save_optimizer_param_state(optimizer, save_path) logger.info(f"Saved optimizer/scheduler state to {save_path}") if dist.is_initialized(): @@ -601,15 +695,47 @@ def _load_training_state( rank = dist.get_rank() if dist.is_initialized() else 0 state_path = adapter_dir / f"training_state_rank{rank}.pt" - if not state_path.exists(): + # Agreed on before the early return: the restore below is collective, so a rank that + # skipped it would leave its peers waiting. + if not _all_ranks_true(state_path.exists()): + if state_path.exists(): + logger.warning(f"{state_path.name} is missing on some ranks; skipping the optimizer restore") return None - # Optimizer state dicts may contain non-tensor objects (e.g. step counts, - # param group metadata), so full unpickling is required here. - training_state = torch.load(state_path, map_location="cpu", weights_only=False) + training_state = None + load_error = None + try: + # Optimizer state dicts may contain non-tensor objects (e.g. step counts, + # param group metadata), so full unpickling is required here. + training_state = torch.load(state_path, map_location="cpu", weights_only=False) + optimizer.load_state_dict(training_state["optimizer"]) + except Exception as error: + load_error = error + _raise_if_any_rank_failed( + load_error, f"Failed to restore optimizer state on at least one rank ({state_path.name})" + ) - optimizer.load_state_dict(training_state["optimizer"]) - logger.info("Restored optimizer state from LoRA checkpoint") + # The restore below is collective, so every rank must reach the same verdict about the files. + entries = _optimizer_param_state_entries(optimizer, adapter_dir) + present = [path.exists() for child, path in entries if _holds_gathered_param_state(child)] + all_present = _all_ranks_true(all(present)) + none_present = _all_ranks_true(not any(present)) + if all_present: + _load_optimizer_param_state(entries) + logger.info("Restored optimizer state from LoRA checkpoint") + elif none_present: + # Checkpoint written before parameter state was saved: the master weights are still + # the ones built at optimizer construction, i.e. from the base model. + logger.warning( + "No optimizer parameter state next to the LoRA adapter; master weights and Adam " + "moments warm-start from the adapter instead of resuming exactly." + ) + optimizer.reload_model_params() + else: + raise RuntimeError( + "Optimizer parameter state is incomplete: some optimizer_param_state_rank*.pt shards are " + "missing. Resume from a checkpoint that has all of them, or remove them all to warm-start." + ) if opt_param_scheduler is not None and training_state.get("opt_param_scheduler") is not None: opt_param_scheduler.load_state_dict(training_state["opt_param_scheduler"]) diff --git a/tests/fast/backends/megatron_utils/test_lora_optimizer_checkpoint.py b/tests/fast/backends/megatron_utils/test_lora_optimizer_checkpoint.py new file mode 100644 index 00000000000..39707e464fa --- /dev/null +++ b/tests/fast/backends/megatron_utils/test_lora_optimizer_checkpoint.py @@ -0,0 +1,90 @@ +"""A DistributedOptimizer keeps its master weights and Adam moments out of ``state_dict()``, +so a LoRA checkpoint must save and restore them separately.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +import torch +from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer + +import miles.backends.megatron_utils.lora_utils as lora_utils + + +class _Child(DistributedOptimizer): + """Stands in for a real DistributedOptimizer, which needs a model and DDP grad buffers.""" + + def __init__(self, *, stub=False, dp_rank=0, numel_unpadded=None): + self.is_stub_optimizer = stub + self.data_parallel_group = SimpleNamespace(rank=lambda: dp_rank) + self.gbuf_ranges = [] if numel_unpadded is None else [{(torch.bfloat16, torch.bfloat16): None}] + self.buffers = [SimpleNamespace(numel_unpadded=numel_unpadded, params=[torch.zeros(1)])] + self.loaded = "not called" + + def save_parameter_state(self, filename): + if self.data_parallel_group.rank() == 0: + torch.save({"master": 1}, filename) + + def load_parameter_state_from_dp_zero(self, state_dict): + self.loaded = state_dict + + +def _write_training_state(directory): + torch.save( + {"iteration": 3, "optimizer": {"step": 3}, "opt_param_scheduler": {"num_steps": 8}}, + directory / "training_state_rank0.pt", + ) + + +def test_parameter_state_round_trips_through_the_data_parallel_root(tmp_path): + root, stub, peer = _Child(dp_rank=0), _Child(stub=True), _Child(dp_rank=1) + optimizer = MagicMock(chained_optimizers=[root, stub, peer]) + + lora_utils._save_optimizer_param_state(optimizer, tmp_path) + assert [path.name for path in sorted(tmp_path.iterdir())] == ["optimizer_param_state_rank0_optimizer0.pt"] + + _write_training_state(tmp_path) + scheduler = MagicMock() + assert lora_utils._load_training_state(tmp_path, optimizer, scheduler) == 3 + + optimizer.load_state_dict.assert_called_once_with({"step": 3}) + assert root.loaded == {"master": 1} + # Every rank must join the scatter, even the ones that read nothing. + assert peer.loaded is None + assert stub.loaded == "not called" + scheduler.load_state_dict.assert_called_once_with({"num_steps": 8}) + + +def test_checkpoint_without_parameter_state_warm_starts_from_the_adapter(tmp_path): + child = _Child(dp_rank=0) + optimizer = MagicMock(chained_optimizers=[child]) + _write_training_state(tmp_path) + + assert lora_utils._load_training_state(tmp_path, optimizer, None) == 3 + # Masters still hold the base model weights, so they have to be refreshed from the adapter. + optimizer.reload_model_params.assert_called_once_with() + assert child.loaded == "not called" + + +def test_partial_parameter_state_is_rejected(tmp_path): + optimizer = MagicMock(chained_optimizers=[_Child(dp_rank=0), _Child(dp_rank=0)]) + _write_training_state(tmp_path) + (tmp_path / "optimizer_param_state_rank0_optimizer1.pt").touch() + + with pytest.raises(RuntimeError, match="Optimizer parameter state is incomplete"): + lora_utils._load_training_state(tmp_path, optimizer, None) + + +def test_parameter_state_from_a_different_model_is_rejected_before_the_scatter(tmp_path): + child = _Child(dp_rank=0, numel_unpadded=64) + optimizer = MagicMock(chained_optimizers=[child]) + _write_training_state(tmp_path) + torch.save( + {0: {(torch.bfloat16, torch.bfloat16): {"numel_unpadded": 32}}}, + tmp_path / "optimizer_param_state_rank0_optimizer0.pt", + ) + + with pytest.raises(RuntimeError, match="Failed to read optimizer parameter state") as failure: + lora_utils._load_training_state(tmp_path, optimizer, None) + assert "does not match the model" in str(failure.value.__cause__) + assert child.loaded == "not called"