From c7d81d478d2e8db6c8cf58aec1d502da225a4a6b Mon Sep 17 00:00:00 2001 From: Zhichen Zeng Date: Fri, 28 Aug 2026 18:16:38 -0700 Subject: [PATCH] feat: dist_muon offloading support in megatron (#2739) --- miles/backends/megatron_utils/model.py | 7 +- miles/utils/arguments.py | 46 +++-- miles_plugins/optimizers/nvme_stream.py | 112 ++++++++++- .../test_qwen3_4B_muon_offload_disk.py | 176 ++++++++++++++++++ tests/fast/optimizers/__init__.py | 0 tests/fast/optimizers/test_nvme_stream.py | 81 ++++++++ 6 files changed, 403 insertions(+), 19 deletions(-) create mode 100644 tests/e2e/megatron/test_qwen3_4B_muon_offload_disk.py create mode 100644 tests/fast/optimizers/__init__.py create mode 100644 tests/fast/optimizers/test_nvme_stream.py diff --git a/miles/backends/megatron_utils/model.py b/miles/backends/megatron_utils/model.py index bd526e76ba..72d142c852 100644 --- a/miles/backends/megatron_utils/model.py +++ b/miles/backends/megatron_utils/model.py @@ -178,6 +178,10 @@ def setup_model_and_optimizer( config.timers = None if _is_muon_optimizer(config.optimizer): + if args.stream_optimizer_state_to_disk: + from miles_plugins.optimizers.nvme_stream import setup_muon_state_on_disk + + setup_muon_state_on_disk(args) if config.muon_split_qkv and "inkling" in (getattr(args, "custom_model_provider_path", None) or ""): if is_first_replica_megatron_main_rank(): logger.info( @@ -201,7 +205,8 @@ def setup_model_and_optimizer( use_gloo_process_groups=args.use_gloo_process_groups, ) - if args.stream_optimizer_state_to_disk: + if args.stream_optimizer_state_to_disk and not _is_muon_optimizer(config.optimizer): + # Muon took the chunked-offloader route above; this store is DistOpt-only. from miles_plugins.optimizers.nvme_stream import setup_optimizer_state_streaming setup_optimizer_state_streaming(args, optimizer) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index f0be5b7a17..d8626cba37 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -199,14 +199,15 @@ def get_miles_extra_args_provider(add_custom_arguments=None): "--stream-optimizer-state-to-disk", action="store_true", help=( - "Stream the fp32 main params and Adam moments through per-bucket files on " - "node-local NVMe during optimizer.step(), bounding GPU residency to one bucket. " - "For when the optimizer state does not fit the GPU *while the step runs*: " - "--offload-train-target=disk cannot help there, because pause/resume happen at " - "phase boundaries and everything is resident again by the time Adam launches. " - "Bit-identical to keeping the state on GPU, at the cost of disk traffic every " - "step. Distinct from --offload-optimizer-states and --optimizer-cpu-offload, " - "and mutually exclusive with both." + "Hold optimizer state in files on node-local NVMe, for when it does not fit the " + "GPU *while the step runs*; --offload-train-target=disk cannot help there.\n" + "adam: streams fp32 main params and moments through per-bucket files, one bucket " + "resident at a time. Requires the distributed optimizer, excludes " + "--offload-optimizer-states and --optimizer-cpu-offload.\n" + "dist_muon: the disk backend for --chunked-optimizer-state-offload, so pass that " + "plus a non-zero --optimizer-state-offload-fraction. --optimizer-cpu-offload is " + "Adam-only. This bounds host residency, not the GPU restore window -- for that " + "set --optimizer-state-offload-chunk-size-mb, which Megatron warns about at 0." ), ) parser.add_argument( @@ -238,7 +239,8 @@ def get_miles_extra_args_provider(add_custom_arguments=None): "its own subdirectory). Should be fast local NVMe (e.g. /scratch); a tmpfs " "mount, which /tmp is on many systems, keeps the data in RAM and defeats both. " "Files are per-process and overwritten in place every step (bounded size); " - "defaults to $SCRATCH/miles_train_offload_." + "defaults to $SCRATCH/miles_train_offload_. Muon's optimizer-state buffers " + "are unlinked once mapped, so their footprint shows in df but not du." ), ) parser.add_argument( @@ -3378,10 +3380,26 @@ def miles_validate_args(args): "process group, so torch.distributed.get_rank() restarts at 0 per cell and two cells " "on one node would share a store directory" ) - assert args.use_distributed_optimizer, "--stream-optimizer-state-to-disk requires the distributed optimizer" - assert ( - args.optimizer == "adam" - ), f"--stream-optimizer-state-to-disk requires --optimizer adam, got {args.optimizer}" + _muon_disk_state = "muon" in (args.optimizer or "").lower() + if _muon_disk_state: + # Megatron's validate_args has not run yet, so gate on the dist_ prefix rather than + # use_layer_wise_distributed_optimizer. + assert args.optimizer.lower().startswith("dist_"), ( + "--stream-optimizer-state-to-disk with Muon requires the layer-wise distributed " + f"optimizer; pass --optimizer dist_muon, got {args.optimizer}" + ) + assert args.chunked_optimizer_state_offload and args.optimizer_state_offload_fraction > 0.0, ( + "--stream-optimizer-state-to-disk with Muon is the disk backend for the chunked " + "offloader; pass --chunked-optimizer-state-offload and a non-zero " + "--optimizer-state-offload-fraction" + ) + else: + assert ( + args.use_distributed_optimizer + ), "--stream-optimizer-state-to-disk requires the distributed optimizer" + assert ( + args.optimizer == "adam" + ), f"--stream-optimizer-state-to-disk requires --optimizer adam, got {args.optimizer}" assert not (args.multi_lora or is_lora_enabled(args)), ( "--stream-optimizer-state-to-disk does not support LoRA: the LoRA checkpoint path " "persists optimizer.state_dict(), which the store leaves empty, and restores the " @@ -3389,7 +3407,7 @@ def miles_validate_args(args): ) assert not args.optimizer_cpu_offload, "--stream-optimizer-state-to-disk excludes --optimizer-cpu-offload" assert ( - not args.offload_optimizer_states + _muon_disk_state or not args.offload_optimizer_states ), "--stream-optimizer-state-to-disk excludes --offload-optimizer-states" assert ( not args.use_precision_aware_optimizer diff --git a/miles_plugins/optimizers/nvme_stream.py b/miles_plugins/optimizers/nvme_stream.py index bbc4e61c54..3efbae080f 100644 --- a/miles_plugins/optimizers/nvme_stream.py +++ b/miles_plugins/optimizers/nvme_stream.py @@ -21,14 +21,21 @@ they reach this class through four methods: Both directory arguments are checkpoint *bases*; the per-rank layout underneath is this file's business, matching the layout of the live scratch directory. + +Muon takes a different route. Its state already rides Megatron's +``ChunkedOptimizerStateOffloader``, whose only tie to host memory is one allocator, so +``setup_muon_state_on_disk`` swaps that allocator for file-backed tensors and leaves the +rest alone. Those buffers are unlinked once mapped, so they show up in ``df``, not ``du``. """ import atexit +import ctypes import errno import json import logging import os import shutil +import tempfile import time from types import MethodType from typing import TYPE_CHECKING, NamedTuple @@ -68,14 +75,23 @@ def _resize(tensor: torch.Tensor, numel: int) -> None: tensor.untyped_storage().resize_(numel * tensor.element_size()) -def _allocate_file(path: str, nbytes: int) -> int: - fd = os.open(path, os.O_RDWR | os.O_CREAT | os.O_CLOEXEC, 0o600) +def _reserve(fd: int, nbytes: int) -> None: + """Reserve blocks up front, so a full filesystem fails here as ENOSPC. + + Sizing a file with ftruncate alone leaves it sparse: the mapping succeeds and the + process dies on SIGBUS at first touch instead, with nothing to point at. + """ try: os.posix_fallocate(fd, 0, nbytes) except OSError as e: if e.errno not in (errno.EOPNOTSUPP, errno.ENOTSUP, errno.EINVAL): raise os.ftruncate(fd, nbytes) + + +def _allocate_file(path: str, nbytes: int) -> int: + fd = os.open(path, os.O_RDWR | os.O_CREAT | os.O_CLOEXEC, 0o600) + _reserve(fd, nbytes) return fd @@ -89,6 +105,43 @@ def _rw_full(op, fd: int, offset: int, buf) -> None: done += n +def _disk_backed_like(tensor: torch.Tensor, directory: str) -> torch.Tensor: + nbytes = max(tensor.numel() * tensor.element_size(), 1) + fd, path = tempfile.mkstemp(dir=directory, suffix=".bin") + try: + _reserve(fd, nbytes) + finally: + os.close(fd) + storage = torch.UntypedStorage.from_file(path, shared=True, nbytes=nbytes) + os.unlink(path) + buffer = torch.empty(0, dtype=tensor.dtype).set_(storage, 0, tensor.shape) + buffer._miles_disk_backed = True + return buffer + + +def _is_disk_backed(tensor: torch.Tensor) -> bool: + return getattr(tensor, "_miles_disk_backed", False) + + +_MS_SYNC = 4 +_libc = ctypes.CDLL(None, use_errno=True) + + +def _flush_mapping(tensor: torch.Tensor) -> int: + """msync one file-backed buffer, returning the bytes it covered. + + Checkpointing calls os.fsync on its own files, which waits on the kernel's writeback + queue -- and our mappings are rewritten every step, so that queue is carrying gigabytes + of our dirty pages by then. Flushing them here keeps that cost attributable and cheap + to repeat: msync over an already-clean mapping returns immediately. + """ + storage = tensor.untyped_storage() + nbytes = storage.nbytes() + if _libc.msync(ctypes.c_void_p(storage.data_ptr()), ctypes.c_size_t(nbytes), _MS_SYNC) != 0: + raise OSError(ctypes.get_errno(), "msync of optimizer state mapping failed") + return nbytes + + def plan_buckets(entries_by_ddp_bucket: dict, limit: int = BUCKET_NUMEL_LIMIT) -> list[list[_Entry]]: planned, current, numel = [], [], 0 for _, entries in sorted(entries_by_ddp_bucket.items(), key=lambda kv: kv[0]): @@ -450,7 +503,7 @@ def setup_optimizer_state_streaming(args, optimizer) -> None: """ from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer - dir_root = os.path.join(args.offload_train_disk_dir, "optimizer_state") + dir_root = _state_dir_root(args) _purge_rank_dir(dir_root) for dist_opt in optimizer.chained_optimizers: assert isinstance( @@ -468,7 +521,57 @@ def setup_optimizer_state_streaming(args, optimizer) -> None: _bind(dist_opt, store) -def _purge_rank_dir(dir_root: str) -> None: +def setup_muon_state_on_disk(args) -> None: + """Back the chunked offloader's host buffers with files, for Muon's optimizer state. + + Must run before the optimizer is built, which is when the offloader is constructed. + """ + from megatron.core.optimizer import optimizer as consuming_module + from megatron.core.optimizer.cpu_offloading import chunked_optimizer_state_offload as defining_module + + base = defining_module.ChunkedOptimizerStateOffloader + if base.__name__ == "DiskOptimizerStateOffloader": + return + rank_dir = _purge_rank_dir(_state_dir_root(args)) + + class DiskOptimizerStateOffloader(base): + state_dir = rank_dir + _disk_bytes = 0 + + def _new_cpu_buffer(self, tensor: torch.Tensor) -> torch.Tensor: # type: ignore[override] + # adopt_cpu_optimizer_state reallocates every non-pinned CPU tensor it finds in + # optimizer.state, and ours never report pinned, so each checkpoint would otherwise + # copy the whole state into fresh mappings. + if _is_disk_backed(tensor): + return tensor + buffer = _disk_backed_like(tensor, self.state_dir) + self._disk_bytes += buffer.numel() * buffer.element_size() + return buffer + + def step(self) -> None: # type: ignore[override] + super().step() + logger.info(f"Muon disk state step: {self._disk_bytes / 1024**3:.2f} GB file-backed") + + def synchronize_for_checkpoint(self) -> None: # type: ignore[override] + # After super(), because it offloads the master weights and so can add mappings. + super().synchronize_for_checkpoint() + flushed = 0 + for state in self._cpu_state.values(): + flushed += sum(_flush_mapping(t) for t in state.values() if _is_disk_backed(t)) + flushed += sum(_flush_mapping(t) for t in self._cpu_master.values() if _is_disk_backed(t)) + logger.info(f"Muon disk state flushed before checkpoint: {flushed / 1024**3:.2f} GB") + + # optimizer.py imported the name directly, so rebinding only the defining module is a no-op. + defining_module.ChunkedOptimizerStateOffloader = DiskOptimizerStateOffloader + consuming_module.ChunkedOptimizerStateOffloader = DiskOptimizerStateOffloader + logger.info(f"Muon optimizer state on disk: buffers backed by files under {rank_dir}") + + +def _state_dir_root(args) -> str: + return os.path.join(args.offload_train_disk_dir, "optimizer_state") + + +def _purge_rank_dir(dir_root: str) -> str: """Drop everything this rank left behind, before any store claims its own path. A store only removes the exact path it is about to use, so state written under a @@ -482,6 +585,7 @@ def _purge_rank_dir(dir_root: str) -> None: rank_dir = os.path.join(dir_root, f"rank{torch.distributed.get_rank():05d}") shutil.rmtree(rank_dir, ignore_errors=True) os.makedirs(rank_dir, exist_ok=True) + return rank_dir def _bind(dist_opt: "DistributedOptimizer", store: NVMeOptimizerStateStore) -> None: diff --git a/tests/e2e/megatron/test_qwen3_4B_muon_offload_disk.py b/tests/e2e/megatron/test_qwen3_4B_muon_offload_disk.py new file mode 100644 index 0000000000..5e87ce6ba5 --- /dev/null +++ b/tests/e2e/megatron/test_qwen3_4B_muon_offload_disk.py @@ -0,0 +1,176 @@ +"""E2E test for dist_muon with its optimizer state on disk. + +The colocate loop from test_qwen3_4B_offload_disk_stream.py with --optimizer dist_muon. +Completing the run is only half the check: if the disk backend silently never engaged, the +run trains just as happily against pinned host memory, so `execute` also asserts every rank +logged real disk-backed steps. + +At this rollout size every sample truncates, so the metric gates below sit at zero and +cannot catch a regression on their own -- as in the Adam test this mirrors. +""" + +import glob +import os + +from tests.ci.ci_register import register_cuda_ci +from tests.ci.metric_history import register_ci_gate + +import miles.utils.external_utils.command_utils as U + +MODEL_NAME = "Qwen3-4B" +MODEL_TYPE = "qwen3-4B" +NUM_GPUS = 4 +OFFLOAD_DIR = "/root/train_offload_muon_disk" + +register_cuda_ci( + est_time=600, + suite="stage-c-4-gpu-h200", + labels=["miles-plugin", "megatron"], +) + +register_ci_gate(metric_key="train/grad_norm") +register_ci_gate(metric_key="train/ppo_kl") +register_ci_gate(metric_key="train/train_rollout_logprob_abs_diff") +register_ci_gate(metric_key="rollout/raw_reward") + + +def prepare(): + U.exec_command_cpu("mkdir -p /root/models /root/datasets") + U.exec_command_cpu(f"hf download Qwen/{MODEL_NAME} --local-dir /root/models/{MODEL_NAME}") + U.hf_download_dataset("zhuzilin/dapo-math-17k") + U.convert_checkpoint(model_name=MODEL_NAME, megatron_model_type=MODEL_TYPE, num_gpus_per_node=NUM_GPUS) + + +def _assert_disk_backed_steps(): + """Every rank must have run optimizer steps against file-backed state.""" + logs = glob.glob("/tmp/ray/session_latest/logs/worker-*") + assert logs, "no Ray worker logs to check for the disk-backed state path" + + backed = set() + for path in logs: + with open(path, errors="ignore") as f: + if any("Muon disk state step:" in line for line in f): + backed.add(path) + + assert ( + len(backed) == NUM_GPUS + ), f"expected {NUM_GPUS} ranks to log disk-backed optimizer steps, saw {len(backed)}: {sorted(backed)}" + print(f"Muon optimizer state was file-backed on {len(backed)} ranks") + + +def _assert_offloaded_to_disk(): + """Every rank must have armed the paused-actor disk offload under its own directory.""" + logs = glob.glob("/tmp/ray/session_latest/logs/worker-*") + assert logs, "no Ray worker logs to check for the disk-offload path" + + armed = set() + for path in logs: + with open(path, errors="ignore") as f: + for line in f: + if "Train disk-offload reclaim armed" in line: + armed.add(line.split("reclaim armed for ")[1].split()[0]) + + expected = {os.path.join(OFFLOAD_DIR, f"cell0_rank{rank}") for rank in range(NUM_GPUS)} + assert armed == expected, f"expected disk offload armed for {sorted(expected)}, saw {sorted(armed)}" + print(f"disk offload armed for {len(armed)} ranks under {OFFLOAD_DIR}") + + +def execute(): + ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/{MODEL_NAME}_torch_dist " + + rollout_args = ( + "--prompt-data /root/datasets/dapo-math-17k/dapo-math-17k.jsonl " + "--input-key prompt " + "--label-key label " + "--apply-chat-template " + "--rollout-shuffle " + "--rm-type deepscaler " + "--num-rollout 2 " + "--rollout-batch-size 4 " + "--n-samples-per-prompt 2 " + "--rollout-max-response-len 256 " + "--rollout-temperature 0.8 " + "--global-batch-size 8 " + "--balance-data " + ) + + perf_args = ( + "--tensor-model-parallel-size 2 " + "--sequence-parallel " + "--pipeline-model-parallel-size 1 " + "--context-parallel-size 1 " + "--recompute-granularity full " + "--recompute-method uniform " + "--recompute-num-layers 1 " + "--use-dynamic-batch-size " + "--max-tokens-per-gpu 2048 " + ) + + grpo_args = ( + "--advantage-estimator grpo " + "--use-kl-loss " + "--kl-loss-coef 0.00 " + "--kl-loss-type low_var_kl " + "--entropy-coef 0.00 " + "--eps-clip 0.2 " + "--eps-clip-high 0.28 " + ) + + optimizer_args = ( + "--optimizer dist_muon " + "--lr 1e-5 " + "--lr-decay-style constant " + "--weight-decay 0.1 " + "--adam-beta1 0.9 " + "--adam-beta2 0.98 " + ) + + # The feature under test: Muon's chunked optimizer state in files rather than + # pinned host memory, on top of paused-actor disk offload. + offload_args = ( + "--offload-train " + "--offload-train-target disk " + f"--offload-train-disk-dir {OFFLOAD_DIR} " + "--offload-train-disk-chunk-mb 64 " + "--chunked-optimizer-state-offload " + "--optimizer-state-offload-fraction 1.0 " + "--stream-optimizer-state-to-disk " + ) + + sglang_args = "--rollout-num-gpus-per-engine 1 " "--sglang-mem-fraction-static 0.6 " + + ci_args = "--ci-test " + + misc_args = ( + "--attention-dropout 0.0 " + "--hidden-dropout 0.0 " + "--accumulate-allreduce-grads-in-fp32 " + "--attention-softmax-in-fp32 " + "--attention-backend flash " + "--actor-num-nodes 1 " + f"--actor-num-gpus-per-node {NUM_GPUS} " + "--colocate " + ) + + train_args = ( + f"{ckpt_args} " + f"{rollout_args} " + f"{optimizer_args} " + f"{grpo_args} " + f"{offload_args} " + f"{U.get_default_wandb_args(__file__)} " + f"{perf_args} " + f"{sglang_args} " + f"{ci_args} " + f"{misc_args} " + ) + + U.execute_train(train_args=train_args, num_gpus_per_node=NUM_GPUS, megatron_model_type=MODEL_TYPE) + + _assert_offloaded_to_disk() + _assert_disk_backed_steps() + + +if __name__ == "__main__": + prepare() + execute() diff --git a/tests/fast/optimizers/__init__.py b/tests/fast/optimizers/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/fast/optimizers/test_nvme_stream.py b/tests/fast/optimizers/test_nvme_stream.py new file mode 100644 index 0000000000..75264f0dc3 --- /dev/null +++ b/tests/fast/optimizers/test_nvme_stream.py @@ -0,0 +1,81 @@ +"""The file-backed buffers behind Muon's disk-resident optimizer state. + +Guards a silent failure: if the allocator stops returning file-backed storage, the offloader +keeps working against pinned host memory while the log still claims otherwise. +""" + +import os + +import pytest +import torch + +from miles_plugins.optimizers import nvme_stream + + +def test_disk_buffer_matches_shape_and_dtype_and_is_not_pinned(tmp_path): + src = torch.randn(64, 32, dtype=torch.float32) + + buf = nvme_stream._disk_backed_like(src, str(tmp_path)) + + assert buf.shape == src.shape + assert buf.dtype == src.dtype + assert buf.device.type == "cpu" + # The inherited offloader picks its sync/async copy path off is_pinned(). + assert not buf.is_pinned() + + +def test_disk_buffer_round_trip_is_bit_exact(tmp_path): + src = torch.randn(128, 64, dtype=torch.float32) + buf = nvme_stream._disk_backed_like(src, str(tmp_path)) + + buf.copy_(src) + out = torch.empty_like(src) + out.copy_(buf) + + assert torch.equal(out, src) + + +def test_disk_buffer_leaves_no_file_behind(tmp_path): + nvme_stream._disk_backed_like(torch.zeros(8), str(tmp_path)) + + # Unlinked at creation, so a killed run leaves no residue. + assert os.listdir(tmp_path) == [] + + +def test_disk_buffer_is_recognized_as_already_managed(tmp_path): + """Megatron's checkpoint adoption reallocates non-pinned CPU state; ours must be exempt.""" + buf = nvme_stream._disk_backed_like(torch.zeros(32, 8), str(tmp_path)) + + assert nvme_stream._is_disk_backed(buf) + assert not nvme_stream._is_disk_backed(torch.zeros(32, 8)) + + +def test_flush_mapping_covers_the_buffer_and_repeats_cheaply(tmp_path): + """Checkpointing fsyncs its own files behind the kernel's writeback of ours.""" + buf = nvme_stream._disk_backed_like(torch.zeros(1024, 256), str(tmp_path)) + nbytes = buf.numel() * buf.element_size() + buf.fill_(1.0) + + assert nvme_stream._flush_mapping(buf) == nbytes + # Already clean, so the repeat is the cheap case the checkpoint hook relies on. + assert nvme_stream._flush_mapping(buf) == nbytes + + +def test_reserve_sizes_the_file(tmp_path): + path = tmp_path / "f.bin" + fd = os.open(str(path), os.O_RDWR | os.O_CREAT, 0o600) + try: + nvme_stream._reserve(fd, 4096) + assert os.fstat(fd).st_size == 4096 + finally: + os.close(fd) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_disk_buffer_preserves_dtype(tmp_path, dtype): + src = torch.zeros(16, 4, dtype=dtype) + + buf = nvme_stream._disk_backed_like(src, str(tmp_path)) + + assert buf.dtype is dtype + assert buf.numel() == src.numel()