feat: dist_muon offloading support in megatron (#2739)

This commit is contained in:
Zhichen Zeng
2026-08-28 18:16:38 -07:00
committed by GitHub
parent dbbab1566a
commit c7d81d478d
6 changed files with 403 additions and 19 deletions
+6 -1
View File
@@ -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)
+32 -14
View File
@@ -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_<uid>."
"defaults to $SCRATCH/miles_train_offload_<uid>. 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
+108 -4
View File
@@ -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:
@@ -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()
View File
+81
View File
@@ -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()