From 51e065bd013337953954474e53a59833cb709ac1 Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Tue, 7 Jul 2026 16:55:39 -0700 Subject: [PATCH 01/12] Add NVMe streaming optimizer state store (v1) fp32 main params and Adam moments live in per-bucket NVMe files and are streamed through the GPU bucket-by-bucket during the optimizer step via per-bucket FusedAdam instances. Bounds step-time GPU residency to one bucket. Checkpointing optimizer state is guarded (not supported yet). --- megatron/core/optimizer/distrib_optimizer.py | 59 ++++-- megatron/core/optimizer/nvme_state_store.py | 206 +++++++++++++++++++ megatron/core/optimizer/optimizer_config.py | 10 + megatron/training/arguments.py | 7 + 4 files changed, 264 insertions(+), 18 deletions(-) create mode 100644 megatron/core/optimizer/nvme_state_store.py diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 8fe58f92bbb..2fb1a351cea 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -537,6 +537,7 @@ def __init__( ) self._state_offloader: Optional[OptimizerStateOffloader] = None + self._nvme_state_store = None # when freezing sub-models we have no real optimizer # but still need a stub DistributedOptimizer class @@ -630,6 +631,15 @@ def __init__( if self.config.offload_optimizer_states: self._state_offloader = OptimizerStateOffloader(self) + if self.config.optimizer_state_nvme_dir is not None: + from megatron.core.optimizer.nvme_state_store import NVMeOptimizerStateStore + + self._nvme_state_store = NVMeOptimizerStateStore( + self, + self.config.optimizer_state_nvme_dir, + self.config.optimizer_state_nvme_chunk_mb, + ) + def _get_model_param_range_map(self, param: torch.nn.Parameter): """ Given a model param, get the index sub-range of the param that this @@ -656,6 +666,11 @@ def state_dict(self): optimizer state (e.g., exp_avg, exp_avg_sq) are stored in a separate checkpoint file by calling 'save_parameter_state()'. """ + if self._nvme_state_store is not None: + raise RuntimeError( + "Checkpointing optimizer state is not supported with " + "--optimizer-state-nvme-dir yet (state lives on NVMe)." + ) inner_state_dict = self.optimizer.state_dict() state_dict = {} @@ -1247,6 +1262,11 @@ def sharded_state_dict( Regular state dict parameters are saved on DP rank 0 and loaded on all ranks. """ + if self._nvme_state_store is not None: + raise RuntimeError( + "Checkpointing optimizer state is not supported with " + "--optimizer-state-nvme-dir yet (state lives on NVMe)." + ) if sharding_type is not None: log_single_rank( logger, @@ -2486,29 +2506,29 @@ def _copy_main_params_to_model_params(self): # Utility method for copying group params. def copy_group_params(shard_main_groups, model_groups): for shard_main_group, model_group in zip(shard_main_groups, model_groups): - for shard_main_param, model_param in zip(shard_main_group, model_group): + self._copy_main_params_to_model_params_for(zip(shard_main_group, model_group)) - param_range_map = self._get_model_param_range_map(model_param) - world_range = param_range_map["gbuf_world_in_bucket"] + # Copy shard groups to model groups. + copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups) + copy_group_params(self.shard_fp32_groups, self.model_fp32_groups) - assert world_range.size == shard_main_param.nelement() + def _copy_main_params_to_model_params_for(self, pairs): + """Copy (shard_main_param, model_param) pairs into the param buffer.""" + for shard_main_param, model_param in pairs: + param_range_map = self._get_model_param_range_map(model_param) + world_range = param_range_map["gbuf_world_in_bucket"] - gbuf_index, _, bucket_id = self.model_param_gbuf_map[model_param] - model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data + assert world_range.size == shard_main_param.nelement() - shard_model_param = model_param_buffer.view(-1)[ - world_range.start : world_range.end - ] + gbuf_index, _, bucket_id = self.model_param_gbuf_map[model_param] + model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data - if is_float8tensor(model_param): - # FP8 params are quantized in the above "quantize_param_shard" function. - continue - else: - shard_model_param.data.copy_(shard_main_param) + shard_model_param = model_param_buffer.view(-1)[world_range.start : world_range.end] - # Copy shard groups to model groups. - copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups) - copy_group_params(self.shard_fp32_groups, self.model_fp32_groups) + if is_float8tensor(model_param): + # FP8 params are quantized in the above "quantize_param_shard" function. + continue + shard_model_param.data.copy_(shard_main_param) def _copy_main_params_to_param_buffer(self): """ @@ -2634,7 +2654,10 @@ def step_with_ready_grads(self) -> bool: """ if self._state_offloader is not None: self._state_offloader.sync_before_step() - update_successful = super().step_with_ready_grads() + if self._nvme_state_store is not None: + update_successful = self._nvme_state_store.step() + else: + update_successful = super().step_with_ready_grads() timers = self.config.timers if timers is not None: diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py new file mode 100644 index 00000000000..130f696101f --- /dev/null +++ b/megatron/core/optimizer/nvme_state_store.py @@ -0,0 +1,206 @@ +import atexit +import logging +import os +import shutil +import time +from typing import TYPE_CHECKING, Dict, List, Tuple + +import torch + +if TYPE_CHECKING: + from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer + +logger = logging.getLogger(__name__) + +_MOMENT_KEYS = ("exp_avg", "exp_avg_sq") + + +class _BucketSpec: + """One DDP bucket's slice of optimizer state and its backing file. + + The file holds three equally-sized segments, [main | exp_avg | exp_avg_sq], + each laid out as the bucket's param shards concatenated in group order. + """ + + def __init__(self, index: int, path: str, entries: List[Tuple[torch.nn.Parameter, torch.Tensor, int]]): + self.index = index + self.path = path + self.entries = entries # (model_param, shard_main_param, master_group_idx) + self.numel = sum(main.numel() for _, main, _ in entries) + offsets = [] + pos = 0 + for _, main, _ in entries: + offsets.append(pos) + pos += main.numel() + self.entry_offsets = offsets + self.fd = -1 + self.adam = None + self.group_master_indices: List[int] = [] + self.main_on_disk = False + self.moments_on_disk = False + + +class NVMeOptimizerStateStore: + """Owns residency and I/O of one DistributedOptimizer's state.""" + + def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int): + self.dist_opt = distrib_optimizer + config = distrib_optimizer.config + + assert not config.use_precision_aware_optimizer, ( + "NVMe state store requires the non-precision-aware optimizer " + "(fp32 main params held by mcore)." + ) + assert not config.optimizer_cpu_offload, "NVMe state store is mutually exclusive with CPU offload." + assert not config.offload_optimizer_states, ( + "NVMe state store is mutually exclusive with --offload-optimizer-states." + ) + assert not distrib_optimizer.ddp_config.use_megatron_fsdp + assert all(len(g) == 0 for g in distrib_optimizer.model_fp32_groups), ( + "NVMe state store only supports pure bf16/fp16 models (no fp32 model params)." + ) + + rank = torch.distributed.get_rank() + instance = distrib_optimizer.distributed_optimizer_instance_id + self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}") + shutil.rmtree(self.dir, ignore_errors=True) + os.makedirs(self.dir, exist_ok=True) + atexit.register(shutil.rmtree, self.dir, ignore_errors=True) + + self._chunk = torch.empty(chunk_mb * 1024 * 1024 // 4, dtype=torch.float32, pin_memory=True) + self._chunk_np = self._chunk.numpy() + + self.specs = self._build_specs() + self._build_bucket_optimizers() + for spec in self.specs: + spec.fd = os.open(spec.path, os.O_RDWR | os.O_CREAT, 0o600) + os.posix_fallocate(spec.fd, 0, 3 * spec.numel * 4) + + total_gb = sum(3 * s.numel * 4 for s in self.specs) / 1024**3 + logger.info( + f"NVMe optimizer state store: {len(self.specs)} buckets, " + f"{total_gb:.1f} GB state at {self.dir}" + ) + + def _build_specs(self) -> List["_BucketSpec"]: + by_bucket: Dict[Tuple, List] = {} + groups = zip( + self.dist_opt.model_float16_groups, self.dist_opt.shard_fp32_from_float16_groups + ) + for group_idx, (model_group, main_group) in enumerate(groups): + for model_param, main_param in zip(model_group, main_group): + assert main_param is not None and main_param.dtype == torch.float32 + key = self.dist_opt.model_param_gbuf_map[model_param] + by_bucket.setdefault(key, []).append((model_param, main_param, group_idx)) + return [ + _BucketSpec(i, os.path.join(self.dir, f"bucket{i:05d}.bin"), entries) + for i, (_, entries) in enumerate(sorted(by_bucket.items(), key=lambda kv: kv[0])) + ] + + def _build_bucket_optimizers(self) -> None: + from megatron.core.optimizer import Adam + + master_groups = self.dist_opt.optimizer.param_groups + for spec in self.specs: + groups = [] + spec.group_master_indices = sorted({gi for _, _, gi in spec.entries}) + for g_idx in spec.group_master_indices: + group = {k: v for k, v in master_groups[g_idx].items() if k != "params"} + group["params"] = [main for _, main, gi in spec.entries if gi == g_idx] + groups.append(group) + spec.adam = Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) + + # ------------------------------------------------------------------ step + + @torch.no_grad() + def step(self) -> bool: + t0 = time.monotonic() + read_bytes = written_bytes = 0 + for spec in self.specs: + read_bytes += self._load_bucket(spec) + self._sync_hyperparams(spec) + spec.adam.step() + self.dist_opt._copy_main_params_to_model_params_for( + (main, model) for model, main, _ in spec.entries + ) + written_bytes += self._store_bucket(spec) + logger.info( + f"NVMe streaming step: {len(self.specs)} buckets, " + f"read {read_bytes / 1024**3:.1f} GB, wrote {written_bytes / 1024**3:.1f} GB " + f"in {time.monotonic() - t0:.1f}s" + ) + return True + + def _sync_hyperparams(self, spec: "_BucketSpec") -> None: + master_groups = self.dist_opt.optimizer.param_groups + for group, g_idx in zip(spec.adam.param_groups, spec.group_master_indices): + group["lr"] = master_groups[g_idx]["lr"] + group["weight_decay"] = master_groups[g_idx]["weight_decay"] + + def _load_bucket(self, spec: "_BucketSpec") -> int: + nbytes = 0 + if spec.main_on_disk: + for tensor, offset in self._segment(spec, "main"): + self._materialize(tensor) + self._stream(spec.fd, offset, tensor, to_disk=False) + nbytes += tensor.numel() * 4 + if spec.moments_on_disk: + for key in _MOMENT_KEYS: + for tensor, offset in self._segment(spec, key): + self._materialize(tensor) + self._stream(spec.fd, offset, tensor, to_disk=False) + nbytes += tensor.numel() * 4 + return nbytes + + def _store_bucket(self, spec: "_BucketSpec") -> int: + nbytes = 0 + for key in ("main",) + _MOMENT_KEYS: + for tensor, offset in self._segment(spec, key): + self._stream(spec.fd, offset, tensor, to_disk=True) + self._release(tensor) + nbytes += tensor.numel() * 4 + spec.main_on_disk = True + spec.moments_on_disk = True + return nbytes + + # ------------------------------------------------------- residency & I/O + + def _segment(self, spec: "_BucketSpec", key: str): + segment_index = ("main",) + _MOMENT_KEYS + base = segment_index.index(key) * spec.numel * 4 + for (_, main, _), entry_offset in zip(spec.entries, spec.entry_offsets): + tensor = main if key == "main" else spec.adam.state[main][key] + yield tensor, base + entry_offset * 4 + + @staticmethod + def _materialize(tensor: torch.Tensor) -> None: + tensor.untyped_storage().resize_(tensor.numel() * tensor.element_size()) + + @staticmethod + def _release(tensor: torch.Tensor) -> None: + tensor.untyped_storage().resize_(0) + + def _stream(self, fd: int, base_offset: int, tensor: torch.Tensor, *, to_disk: bool) -> None: + flat = tensor.view(-1) + chunk_numel = self._chunk.numel() + pos = 0 + while pos < flat.numel(): + n = min(chunk_numel, flat.numel() - pos) + byte_offset = base_offset + pos * 4 + if to_disk: + self._chunk[:n].copy_(flat[pos : pos + n]) + self._rw_full(os.pwritev, fd, byte_offset, self._chunk_np[:n]) + else: + self._rw_full(os.preadv, fd, byte_offset, self._chunk_np[:n]) + flat[pos : pos + n].copy_(self._chunk[:n]) + pos += n + + @staticmethod + def _rw_full(op, fd: int, offset: int, array) -> None: + mv = memoryview(array).cast("B") + done = 0 + while done < len(mv): + n = op(fd, [mv[done:]], offset + done) + if n <= 0: + raise IOError(f"short {op.__name__} ({n}) on optimizer state file at offset {offset + done}") + done += n diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 0f7081f4fc1..c4ca7c95b47 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -333,6 +333,16 @@ class OptimizerConfig: low_memory_resume: bool = False """If True, allocate optimizer states on CPU during checkpoint loading to prevent GPU OOM.""" + optimizer_state_nvme_dir: Optional[str] = None + """ + If set, fp32 main params and Adam moments live in per-bucket files under this + node-local directory and are streamed through the GPU bucket-by-bucket during + the optimizer step, bounding GPU residency to one bucket. + """ + + optimizer_state_nvme_chunk_mb: int = 256 + """Pinned staging chunk size for NVMe optimizer state streaming.""" + ################ # Miscellaneous ################ diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 28fea46195d..9c7d2e3fa0b 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2110,6 +2110,13 @@ def _add_training_args(parser): 'Only support TE FusedAdam optimizer.' 'Note that this still uses pure GPU optimizer instead of ' 'HybridDeviceOptimizer for --optimizer-cpu-offload.') + group.add_argument('--optimizer-state-nvme-dir', type=str, default=None, + help='Stream fp32 main params and Adam moments through per-bucket ' + 'files under this node-local directory during the optimizer step, ' + 'bounding GPU residency to one bucket. Checkpointing optimizer ' + 'state is not supported yet.') + group.add_argument('--optimizer-state-nvme-chunk-mb', type=int, default=256, + help='Pinned staging chunk size for NVMe optimizer state streaming.') group.add_argument('--dataloader-type', type=str, default=None, choices=['single', 'cyclic', 'external'], help='Single pass vs multiple pass data loader') From 13a21d94a781705f5d8ec0bb2d8a42dbbcfffba6 Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Wed, 8 Jul 2026 23:26:24 -0700 Subject: [PATCH 02/12] Harden NVMe state store for large-model DP1 - Evict fp32 main params to NVMe at store init: with lazy eviction they stay GPU-resident through the entire first rollout + forward/backward and OOM tight configs before the first step can flush them. - Route _copy_model_params_to_main_params through the store (materialize, copy, flush, release): checkpoint load refreshes main params, which otherwise writes into evicted storage (illegal memory access). - Chunk state specs to a fixed 200M-numel budget: DDP bucketing leaves an entire expert buffer as a single bucket when expert-DP is 1, so bucket-granular streaming materializes ~90 GiB at once on those ranks. Validated on GLM-5.2 744B, 8x GB300 (TP8/PP4/DP1/EP8): two full RL steps, flat 142 GiB watermark through the 120-bucket step loop, actor_train 287s/216s (16-node baseline ~170s), train_rollout_logprob_abs_diff in the reference band. --- megatron/core/optimizer/distrib_optimizer.py | 9 ++++++ megatron/core/optimizer/nvme_state_store.py | 32 +++++++++++++++++++- 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 2fb1a351cea..d9f22e97ab5 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -2591,6 +2591,15 @@ def _build_model_param_to_state_dict_param_map(self, state_dict): return model_param_to_state_dict_param_map def _copy_model_params_to_main_params(self, state_dict=None): + if self._nvme_state_store is not None: + self._nvme_state_store.refresh_main_from_model_params( + lambda: self._copy_model_params_to_main_params_impl(state_dict) + ) + return + self._copy_model_params_to_main_params_impl(state_dict) + + @torch.no_grad() + def _copy_model_params_to_main_params_impl(self, state_dict=None): """ Copy model params to main params. diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py index 130f696101f..b92e533a667 100644 --- a/megatron/core/optimizer/nvme_state_store.py +++ b/megatron/core/optimizer/nvme_state_store.py @@ -76,6 +76,12 @@ def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chu spec.fd = os.open(spec.path, os.O_RDWR | os.O_CREAT, 0o600) os.posix_fallocate(spec.fd, 0, 3 * spec.numel * 4) + for spec in self.specs: + for tensor, offset in self._segment(spec, "main"): + self._stream(spec.fd, offset, tensor, to_disk=True) + self._release(tensor) + spec.main_on_disk = True + total_gb = sum(3 * s.numel * 4 for s in self.specs) / 1024**3 logger.info( f"NVMe optimizer state store: {len(self.specs)} buckets, " @@ -92,9 +98,21 @@ def _build_specs(self) -> List["_BucketSpec"]: assert main_param is not None and main_param.dtype == torch.float32 key = self.dist_opt.model_param_gbuf_map[model_param] by_bucket.setdefault(key, []).append((model_param, main_param, group_idx)) + limit = 200_000_000 + chunked = [] + for _, entries in sorted(by_bucket.items(), key=lambda kv: kv[0]): + cur, cur_numel = [], 0 + for entry in entries: + cur.append(entry) + cur_numel += entry[1].numel() + if cur_numel >= limit: + chunked.append(cur) + cur, cur_numel = [], 0 + if cur: + chunked.append(cur) return [ _BucketSpec(i, os.path.join(self.dir, f"bucket{i:05d}.bin"), entries) - for i, (_, entries) in enumerate(sorted(by_bucket.items(), key=lambda kv: kv[0])) + for i, entries in enumerate(chunked) ] def _build_bucket_optimizers(self) -> None: @@ -112,6 +130,18 @@ def _build_bucket_optimizers(self) -> None: # ------------------------------------------------------------------ step + @torch.no_grad() + def refresh_main_from_model_params(self, copy_fn) -> None: + for spec in self.specs: + for tensor, _ in self._segment(spec, "main"): + self._materialize(tensor) + copy_fn() + for spec in self.specs: + for tensor, offset in self._segment(spec, "main"): + self._stream(spec.fd, offset, tensor, to_disk=True) + self._release(tensor) + spec.main_on_disk = True + @torch.no_grad() def step(self) -> bool: t0 = time.monotonic() From ff23ee27a96d4d798f841278f7521cb4dde90118 Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Thu, 9 Jul 2026 11:29:06 -0700 Subject: [PATCH 03/12] NVMe store checkpointing: save/load bucket files with manifest (same-topology resume) --- megatron/core/optimizer/distrib_optimizer.py | 13 +++--- megatron/core/optimizer/nvme_state_store.py | 47 ++++++++++++++++++++ megatron/training/checkpointing.py | 30 +++++++++++++ 3 files changed, 82 insertions(+), 8 deletions(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index d9f22e97ab5..2c45f178d5d 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -667,10 +667,7 @@ def state_dict(self): checkpoint file by calling 'save_parameter_state()'. """ if self._nvme_state_store is not None: - raise RuntimeError( - "Checkpointing optimizer state is not supported with " - "--optimizer-state-nvme-dir yet (state lives on NVMe)." - ) + return {"nvme_state_store": True} inner_state_dict = self.optimizer.state_dict() state_dict = {} @@ -756,6 +753,9 @@ def load_state_dict(self, state_dict): - state_order : The index of a parameter within the shared parameter list. """ + if self._nvme_state_store is not None: + return + if self.ddp_config.use_megatron_fsdp: if "param_to_group_meta" in state_dict: state_dict["param_groups"] = self._param2group_meta_to_param_groups( @@ -1263,10 +1263,7 @@ def sharded_state_dict( Regular state dict parameters are saved on DP rank 0 and loaded on all ranks. """ if self._nvme_state_store is not None: - raise RuntimeError( - "Checkpointing optimizer state is not supported with " - "--optimizer-state-nvme-dir yet (state lives on NVMe)." - ) + return {} if sharding_type is not None: log_single_rank( logger, diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py index b92e533a667..b35b8d90fc3 100644 --- a/megatron/core/optimizer/nvme_state_store.py +++ b/megatron/core/optimizer/nvme_state_store.py @@ -1,4 +1,5 @@ import atexit +import json import logging import os import shutil @@ -130,6 +131,52 @@ def _build_bucket_optimizers(self) -> None: # ------------------------------------------------------------------ step + @torch.no_grad() + def save_to(self, dirpath: str) -> None: + os.makedirs(dirpath, exist_ok=True) + manifest = { + "buckets": [ + { + "numel": spec.numel, + "entry_numels": [main.numel() for _, main, _ in spec.entries], + "steps": [g.get("step", 0) for g in spec.adam.param_groups], + "file": os.path.basename(spec.path), + } + for spec in self.specs + ] + } + for spec in self.specs: + shutil.copyfile(spec.path, os.path.join(dirpath, os.path.basename(spec.path))) + with open(os.path.join(dirpath, "manifest.json"), "w") as f: + json.dump(manifest, f) + logger.info(f"NVMe optimizer state saved: {len(self.specs)} buckets -> {dirpath}") + + @torch.no_grad() + def load_from(self, dirpath: str) -> None: + with open(os.path.join(dirpath, "manifest.json")) as f: + manifest = json.load(f) + assert len(manifest["buckets"]) == len(self.specs), ( + f"NVMe state layout mismatch: checkpoint has {len(manifest['buckets'])} buckets, " + f"current topology builds {len(self.specs)} (same-topology resume only)" + ) + for spec, meta in zip(self.specs, manifest["buckets"]): + assert meta["numel"] == spec.numel + assert meta["entry_numels"] == [main.numel() for _, main, _ in spec.entries] + shutil.copyfile(os.path.join(dirpath, meta["file"]), spec.path) + for group, step in zip(spec.adam.param_groups, meta["steps"]): + if step: + group["step"] = step + for _, main, _ in spec.entries: + state = spec.adam.state.setdefault(main, {}) + for key in _MOMENT_KEYS: + if key not in state: + t = torch.empty_like(main) + t.untyped_storage().resize_(0) + state[key] = t + spec.main_on_disk = True + spec.moments_on_disk = True + logger.info(f"NVMe optimizer state loaded: {len(self.specs)} buckets <- {dirpath}") + @torch.no_grad() def refresh_main_from_model_params(self, copy_fn) -> None: for spec in self.specs: diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 7c87eca191a..34075bac5f7 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -468,6 +468,19 @@ def save_grads(save_dir, state_dict, iteration, grad_label): f"from iteration {iteration:7d}") +def _iter_nvme_state_stores(optimizer): + for opt in getattr(optimizer, "chained_optimizers", None) or [optimizer]: + store = getattr(opt, "_nvme_state_store", None) + if store is not None: + yield store + + +def _nvme_state_checkpoint_dir(checkpoint_name, store): + rank = torch.distributed.get_rank() + instance = store.dist_opt.distributed_optimizer_instance_id + return os.path.join(checkpoint_name, "nvme_opt_state", f"rank{rank:04d}_opt{instance}") + + def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floating_point_operations_so_far, checkpointing_context=None, pipeline_rank=None, expert_rank=None, tensor_rank=None, pipeline_parallel=None, expert_parallel=None, non_persistent_ckpt=False, train_data_iterator=None, preprocess_common_state_dict_fn = None, release=False, tp_group: Optional[torch.distributed.ProcessGroup] = None, pp_group: Optional[torch.distributed.ProcessGroup] = None, dp_cp_group: Optional[torch.distributed.ProcessGroup] = None): @@ -570,6 +583,12 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati if not optimizer.is_stub_optimizer: optimizer.save_state_dict_to_file(optim_checkpoint_name) + # NVMe-streamed optimizer state (--optimizer-state-nvme-dir): the flat + # bucket files on node-local scratch are the state; copy them per rank. + if not args.no_save_optim and optimizer is not None: + for store in _iter_nvme_state_stores(optimizer): + store.save_to(_nvme_state_checkpoint_dir(checkpoint_name, store)) + async_save_request = None if args.async_save: if ckpt_type == CheckpointType.LEGACY: @@ -1829,6 +1848,17 @@ def load_model_state_dict(module, state_dict, strict: bool): else: optimizer.reload_model_params() + # NVMe-streamed optimizer state: restore the bucket files after + # reload_model_params so the checkpointed fp32 main wins over the + # bf16-recast refresh. + if optimizer is not None and not release and not args.finetune and not args.no_load_optim: + for store in _iter_nvme_state_stores(optimizer): + nvme_dir = _nvme_state_checkpoint_dir(checkpoint_name, store) + if os.path.isdir(nvme_dir): + store.load_from(nvme_dir) + else: + print_rank_0(f" no NVMe optimizer state at {nvme_dir}; starting fresh") + # rerun state if not ignore_rerun_state: try: From 34d73a3c2e5adfe58b9a167b579060107c824d6a Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Thu, 9 Jul 2026 12:35:37 -0700 Subject: [PATCH 04/12] Disambiguate store directories: chained optimizer instances can share instance_id --- megatron/core/optimizer/nvme_state_store.py | 10 +++++++++- megatron/training/checkpointing.py | 4 +++- 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py index b35b8d90fc3..e21acb8c961 100644 --- a/megatron/core/optimizer/nvme_state_store.py +++ b/megatron/core/optimizer/nvme_state_store.py @@ -44,8 +44,16 @@ def __init__(self, index: int, path: str, entries: List[Tuple[torch.nn.Parameter class NVMeOptimizerStateStore: """Owns residency and I/O of one DistributedOptimizer's state.""" + # ChainedOptimizer members (dense/expert) can share + # distributed_optimizer_instance_id, so a per-process counter keeps their + # store directories distinct. Construction order is deterministic, which + # also keeps checkpoint directory names stable across runs. + _next_uid = 0 + def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int): self.dist_opt = distrib_optimizer + self.uid = NVMeOptimizerStateStore._next_uid + NVMeOptimizerStateStore._next_uid += 1 config = distrib_optimizer.config assert not config.use_precision_aware_optimizer, ( @@ -63,7 +71,7 @@ def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chu rank = torch.distributed.get_rank() instance = distrib_optimizer.distributed_optimizer_instance_id - self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}") + self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}_{self.uid}") shutil.rmtree(self.dir, ignore_errors=True) os.makedirs(self.dir, exist_ok=True) atexit.register(shutil.rmtree, self.dir, ignore_errors=True) diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 34075bac5f7..a6464fa50df 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -478,7 +478,9 @@ def _iter_nvme_state_stores(optimizer): def _nvme_state_checkpoint_dir(checkpoint_name, store): rank = torch.distributed.get_rank() instance = store.dist_opt.distributed_optimizer_instance_id - return os.path.join(checkpoint_name, "nvme_opt_state", f"rank{rank:04d}_opt{instance}") + return os.path.join( + checkpoint_name, "nvme_opt_state", f"rank{rank:04d}_opt{instance}_{store.uid}" + ) def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floating_point_operations_so_far, From 17ca37315125236b1927ac7b45ef2973b2989aac Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Mon, 20 Jul 2026 20:25:41 -0700 Subject: [PATCH 05/12] Trim redundant NVMe optimizer-state comments --- megatron/core/optimizer/distrib_optimizer.py | 1 - megatron/core/optimizer/nvme_state_store.py | 6 ++---- megatron/core/optimizer/optimizer_config.py | 7 ++----- megatron/training/checkpointing.py | 5 ----- 4 files changed, 4 insertions(+), 15 deletions(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 2c45f178d5d..77c1410f01d 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -2505,7 +2505,6 @@ def copy_group_params(shard_main_groups, model_groups): for shard_main_group, model_group in zip(shard_main_groups, model_groups): self._copy_main_params_to_model_params_for(zip(shard_main_group, model_group)) - # Copy shard groups to model groups. copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups) copy_group_params(self.shard_fp32_groups, self.model_fp32_groups) diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py index e21acb8c961..fb9e8f90dc0 100644 --- a/megatron/core/optimizer/nvme_state_store.py +++ b/megatron/core/optimizer/nvme_state_store.py @@ -44,10 +44,8 @@ def __init__(self, index: int, path: str, entries: List[Tuple[torch.nn.Parameter class NVMeOptimizerStateStore: """Owns residency and I/O of one DistributedOptimizer's state.""" - # ChainedOptimizer members (dense/expert) can share - # distributed_optimizer_instance_id, so a per-process counter keeps their - # store directories distinct. Construction order is deterministic, which - # also keeps checkpoint directory names stable across runs. + # ChainedOptimizer's dense/expert members can share distributed_optimizer_instance_id; + # this per-process counter keeps their store directories distinct. _next_uid = 0 def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int): diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index c4ca7c95b47..a62a9a98173 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -334,11 +334,8 @@ class OptimizerConfig: """If True, allocate optimizer states on CPU during checkpoint loading to prevent GPU OOM.""" optimizer_state_nvme_dir: Optional[str] = None - """ - If set, fp32 main params and Adam moments live in per-bucket files under this - node-local directory and are streamed through the GPU bucket-by-bucket during - the optimizer step, bounding GPU residency to one bucket. - """ + """If set, stream fp32 main params and Adam moments through per-bucket files under this + node-local directory during the optimizer step, bounding GPU residency to one bucket.""" optimizer_state_nvme_chunk_mb: int = 256 """Pinned staging chunk size for NVMe optimizer state streaming.""" diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index a6464fa50df..0ea6101d0cb 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -585,8 +585,6 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati if not optimizer.is_stub_optimizer: optimizer.save_state_dict_to_file(optim_checkpoint_name) - # NVMe-streamed optimizer state (--optimizer-state-nvme-dir): the flat - # bucket files on node-local scratch are the state; copy them per rank. if not args.no_save_optim and optimizer is not None: for store in _iter_nvme_state_stores(optimizer): store.save_to(_nvme_state_checkpoint_dir(checkpoint_name, store)) @@ -1850,9 +1848,6 @@ def load_model_state_dict(module, state_dict, strict: bool): else: optimizer.reload_model_params() - # NVMe-streamed optimizer state: restore the bucket files after - # reload_model_params so the checkpointed fp32 main wins over the - # bf16-recast refresh. if optimizer is not None and not release and not args.finetune and not args.no_load_optim: for store in _iter_nvme_state_stores(optimizer): nvme_dir = _nvme_state_checkpoint_dir(checkpoint_name, store) From c8b5785b9126265a7b631cbcc00bfc84cf3376b9 Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Mon, 20 Jul 2026 21:34:29 -0700 Subject: [PATCH 06/12] Support models with native-fp32 params (router expert_bias, GDN A_log) Route them through a small always-resident Adam instead of hard-rejecting any model that has them; NVMe-stream only the bf16/fp16 buckets, which is where the actual memory pressure is. --- megatron/core/optimizer/nvme_state_store.py | 60 +++++++++++++++++++-- 1 file changed, 57 insertions(+), 3 deletions(-) diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py index fb9e8f90dc0..ab2a14a157a 100644 --- a/megatron/core/optimizer/nvme_state_store.py +++ b/megatron/core/optimizer/nvme_state_store.py @@ -15,6 +15,11 @@ _MOMENT_KEYS = ("exp_avg", "exp_avg_sq") +# Native-fp32 model params (e.g. router expert_bias, GDN/Mamba A_log) stay GPU-resident +# instead of being NVMe-streamed -- warn if they add up to more than this, since the +# assumption that they're negligible in size no longer holds. +_FP32_RESIDENT_WARN_MB = 256 + class _BucketSpec: """One DDP bucket's slice of optimizer state and its backing file. @@ -63,9 +68,6 @@ def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chu "NVMe state store is mutually exclusive with --offload-optimizer-states." ) assert not distrib_optimizer.ddp_config.use_megatron_fsdp - assert all(len(g) == 0 for g in distrib_optimizer.model_fp32_groups), ( - "NVMe state store only supports pure bf16/fp16 models (no fp32 model params)." - ) rank = torch.distributed.get_rank() instance = distrib_optimizer.distributed_optimizer_instance_id @@ -79,6 +81,7 @@ def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chu self.specs = self._build_specs() self._build_bucket_optimizers() + self._build_fp32_optimizer() for spec in self.specs: spec.fd = os.open(spec.path, os.O_RDWR | os.O_CREAT, 0o600) os.posix_fallocate(spec.fd, 0, 3 * spec.numel * 4) @@ -135,6 +138,43 @@ def _build_bucket_optimizers(self) -> None: groups.append(group) spec.adam = Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) + def _build_fp32_optimizer(self) -> None: + """Step native-fp32 model params (router expert_bias, GDN/Mamba A_log, ...) via a + small always-resident Adam instead of NVMe-streaming them. + + ``shard_fp32_groups`` entries are views into the model params' own storage (see + DistributedOptimizer._build_model_and_main_param_groups), so stepping them updates + the model directly -- no copy-back needed, unlike the bf16 bucket path. + """ + from megatron.core.optimizer import Adam + + master_groups = self.dist_opt.optimizer.param_groups + self._fp32_group_indices: List[int] = [] + groups = [] + total_bytes = 0 + for group_idx, (model_group, shard_group) in enumerate( + zip(self.dist_opt.model_fp32_groups, self.dist_opt.shard_fp32_groups) + ): + if not model_group: + continue + group = {k: v for k, v in master_groups[group_idx].items() if k != "params"} + group["params"] = list(shard_group) + groups.append(group) + self._fp32_group_indices.append(group_idx) + total_bytes += sum(p.numel() * p.element_size() for p in shard_group) + + if not groups: + self._fp32_adam = None + return + + total_mb = total_bytes / 1024**2 + log = logger.warning if total_mb > _FP32_RESIDENT_WARN_MB else logger.info + log( + f"NVMe optimizer state store: {total_mb:.1f} MB of native-fp32 model params " + f"stay GPU-resident (not NVMe-managed)." + ) + self._fp32_adam = Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) + # ------------------------------------------------------------------ step @torch.no_grad() @@ -153,6 +193,8 @@ def save_to(self, dirpath: str) -> None: } for spec in self.specs: shutil.copyfile(spec.path, os.path.join(dirpath, os.path.basename(spec.path))) + if self._fp32_adam is not None: + torch.save(self._fp32_adam.state_dict(), os.path.join(dirpath, "fp32_resident_optimizer.pt")) with open(os.path.join(dirpath, "manifest.json"), "w") as f: json.dump(manifest, f) logger.info(f"NVMe optimizer state saved: {len(self.specs)} buckets -> {dirpath}") @@ -181,6 +223,9 @@ def load_from(self, dirpath: str) -> None: state[key] = t spec.main_on_disk = True spec.moments_on_disk = True + fp32_state_path = os.path.join(dirpath, "fp32_resident_optimizer.pt") + if self._fp32_adam is not None and os.path.isfile(fp32_state_path): + self._fp32_adam.load_state_dict(torch.load(fp32_state_path)) logger.info(f"NVMe optimizer state loaded: {len(self.specs)} buckets <- {dirpath}") @torch.no_grad() @@ -207,6 +252,9 @@ def step(self) -> bool: (main, model) for model, main, _ in spec.entries ) written_bytes += self._store_bucket(spec) + if self._fp32_adam is not None: + self._sync_fp32_hyperparams() + self._fp32_adam.step() logger.info( f"NVMe streaming step: {len(self.specs)} buckets, " f"read {read_bytes / 1024**3:.1f} GB, wrote {written_bytes / 1024**3:.1f} GB " @@ -220,6 +268,12 @@ def _sync_hyperparams(self, spec: "_BucketSpec") -> None: group["lr"] = master_groups[g_idx]["lr"] group["weight_decay"] = master_groups[g_idx]["weight_decay"] + def _sync_fp32_hyperparams(self) -> None: + master_groups = self.dist_opt.optimizer.param_groups + for group, g_idx in zip(self._fp32_adam.param_groups, self._fp32_group_indices): + group["lr"] = master_groups[g_idx]["lr"] + group["weight_decay"] = master_groups[g_idx]["weight_decay"] + def _load_bucket(self, spec: "_BucketSpec") -> int: nbytes = 0 if spec.main_on_disk: From 644d83d90a68b76b9a40321833d30bd32eceaa36 Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Wed, 22 Jul 2026 00:22:36 -0700 Subject: [PATCH 07/12] Move NVMe optimizer-state store to miles side Per review: the implementation now lives in miles (miles/backends/megatron_utils/nvme_state_store.py) and is imported lazily at construction, same pattern as the R3 routing-replay hook in moe/router.py. Megatron keeps only the thin config fields and the duck-typed routing hooks in DistributedOptimizer/checkpointing; miles owns the store and can unit-test it. --- megatron/core/optimizer/nvme_state_store.py | 343 -------------------- 1 file changed, 343 deletions(-) delete mode 100644 megatron/core/optimizer/nvme_state_store.py diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py deleted file mode 100644 index ab2a14a157a..00000000000 --- a/megatron/core/optimizer/nvme_state_store.py +++ /dev/null @@ -1,343 +0,0 @@ -import atexit -import json -import logging -import os -import shutil -import time -from typing import TYPE_CHECKING, Dict, List, Tuple - -import torch - -if TYPE_CHECKING: - from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer - -logger = logging.getLogger(__name__) - -_MOMENT_KEYS = ("exp_avg", "exp_avg_sq") - -# Native-fp32 model params (e.g. router expert_bias, GDN/Mamba A_log) stay GPU-resident -# instead of being NVMe-streamed -- warn if they add up to more than this, since the -# assumption that they're negligible in size no longer holds. -_FP32_RESIDENT_WARN_MB = 256 - - -class _BucketSpec: - """One DDP bucket's slice of optimizer state and its backing file. - - The file holds three equally-sized segments, [main | exp_avg | exp_avg_sq], - each laid out as the bucket's param shards concatenated in group order. - """ - - def __init__(self, index: int, path: str, entries: List[Tuple[torch.nn.Parameter, torch.Tensor, int]]): - self.index = index - self.path = path - self.entries = entries # (model_param, shard_main_param, master_group_idx) - self.numel = sum(main.numel() for _, main, _ in entries) - offsets = [] - pos = 0 - for _, main, _ in entries: - offsets.append(pos) - pos += main.numel() - self.entry_offsets = offsets - self.fd = -1 - self.adam = None - self.group_master_indices: List[int] = [] - self.main_on_disk = False - self.moments_on_disk = False - - -class NVMeOptimizerStateStore: - """Owns residency and I/O of one DistributedOptimizer's state.""" - - # ChainedOptimizer's dense/expert members can share distributed_optimizer_instance_id; - # this per-process counter keeps their store directories distinct. - _next_uid = 0 - - def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int): - self.dist_opt = distrib_optimizer - self.uid = NVMeOptimizerStateStore._next_uid - NVMeOptimizerStateStore._next_uid += 1 - config = distrib_optimizer.config - - assert not config.use_precision_aware_optimizer, ( - "NVMe state store requires the non-precision-aware optimizer " - "(fp32 main params held by mcore)." - ) - assert not config.optimizer_cpu_offload, "NVMe state store is mutually exclusive with CPU offload." - assert not config.offload_optimizer_states, ( - "NVMe state store is mutually exclusive with --offload-optimizer-states." - ) - assert not distrib_optimizer.ddp_config.use_megatron_fsdp - - rank = torch.distributed.get_rank() - instance = distrib_optimizer.distributed_optimizer_instance_id - self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}_{self.uid}") - shutil.rmtree(self.dir, ignore_errors=True) - os.makedirs(self.dir, exist_ok=True) - atexit.register(shutil.rmtree, self.dir, ignore_errors=True) - - self._chunk = torch.empty(chunk_mb * 1024 * 1024 // 4, dtype=torch.float32, pin_memory=True) - self._chunk_np = self._chunk.numpy() - - self.specs = self._build_specs() - self._build_bucket_optimizers() - self._build_fp32_optimizer() - for spec in self.specs: - spec.fd = os.open(spec.path, os.O_RDWR | os.O_CREAT, 0o600) - os.posix_fallocate(spec.fd, 0, 3 * spec.numel * 4) - - for spec in self.specs: - for tensor, offset in self._segment(spec, "main"): - self._stream(spec.fd, offset, tensor, to_disk=True) - self._release(tensor) - spec.main_on_disk = True - - total_gb = sum(3 * s.numel * 4 for s in self.specs) / 1024**3 - logger.info( - f"NVMe optimizer state store: {len(self.specs)} buckets, " - f"{total_gb:.1f} GB state at {self.dir}" - ) - - def _build_specs(self) -> List["_BucketSpec"]: - by_bucket: Dict[Tuple, List] = {} - groups = zip( - self.dist_opt.model_float16_groups, self.dist_opt.shard_fp32_from_float16_groups - ) - for group_idx, (model_group, main_group) in enumerate(groups): - for model_param, main_param in zip(model_group, main_group): - assert main_param is not None and main_param.dtype == torch.float32 - key = self.dist_opt.model_param_gbuf_map[model_param] - by_bucket.setdefault(key, []).append((model_param, main_param, group_idx)) - limit = 200_000_000 - chunked = [] - for _, entries in sorted(by_bucket.items(), key=lambda kv: kv[0]): - cur, cur_numel = [], 0 - for entry in entries: - cur.append(entry) - cur_numel += entry[1].numel() - if cur_numel >= limit: - chunked.append(cur) - cur, cur_numel = [], 0 - if cur: - chunked.append(cur) - return [ - _BucketSpec(i, os.path.join(self.dir, f"bucket{i:05d}.bin"), entries) - for i, entries in enumerate(chunked) - ] - - def _build_bucket_optimizers(self) -> None: - from megatron.core.optimizer import Adam - - master_groups = self.dist_opt.optimizer.param_groups - for spec in self.specs: - groups = [] - spec.group_master_indices = sorted({gi for _, _, gi in spec.entries}) - for g_idx in spec.group_master_indices: - group = {k: v for k, v in master_groups[g_idx].items() if k != "params"} - group["params"] = [main for _, main, gi in spec.entries if gi == g_idx] - groups.append(group) - spec.adam = Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) - - def _build_fp32_optimizer(self) -> None: - """Step native-fp32 model params (router expert_bias, GDN/Mamba A_log, ...) via a - small always-resident Adam instead of NVMe-streaming them. - - ``shard_fp32_groups`` entries are views into the model params' own storage (see - DistributedOptimizer._build_model_and_main_param_groups), so stepping them updates - the model directly -- no copy-back needed, unlike the bf16 bucket path. - """ - from megatron.core.optimizer import Adam - - master_groups = self.dist_opt.optimizer.param_groups - self._fp32_group_indices: List[int] = [] - groups = [] - total_bytes = 0 - for group_idx, (model_group, shard_group) in enumerate( - zip(self.dist_opt.model_fp32_groups, self.dist_opt.shard_fp32_groups) - ): - if not model_group: - continue - group = {k: v for k, v in master_groups[group_idx].items() if k != "params"} - group["params"] = list(shard_group) - groups.append(group) - self._fp32_group_indices.append(group_idx) - total_bytes += sum(p.numel() * p.element_size() for p in shard_group) - - if not groups: - self._fp32_adam = None - return - - total_mb = total_bytes / 1024**2 - log = logger.warning if total_mb > _FP32_RESIDENT_WARN_MB else logger.info - log( - f"NVMe optimizer state store: {total_mb:.1f} MB of native-fp32 model params " - f"stay GPU-resident (not NVMe-managed)." - ) - self._fp32_adam = Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) - - # ------------------------------------------------------------------ step - - @torch.no_grad() - def save_to(self, dirpath: str) -> None: - os.makedirs(dirpath, exist_ok=True) - manifest = { - "buckets": [ - { - "numel": spec.numel, - "entry_numels": [main.numel() for _, main, _ in spec.entries], - "steps": [g.get("step", 0) for g in spec.adam.param_groups], - "file": os.path.basename(spec.path), - } - for spec in self.specs - ] - } - for spec in self.specs: - shutil.copyfile(spec.path, os.path.join(dirpath, os.path.basename(spec.path))) - if self._fp32_adam is not None: - torch.save(self._fp32_adam.state_dict(), os.path.join(dirpath, "fp32_resident_optimizer.pt")) - with open(os.path.join(dirpath, "manifest.json"), "w") as f: - json.dump(manifest, f) - logger.info(f"NVMe optimizer state saved: {len(self.specs)} buckets -> {dirpath}") - - @torch.no_grad() - def load_from(self, dirpath: str) -> None: - with open(os.path.join(dirpath, "manifest.json")) as f: - manifest = json.load(f) - assert len(manifest["buckets"]) == len(self.specs), ( - f"NVMe state layout mismatch: checkpoint has {len(manifest['buckets'])} buckets, " - f"current topology builds {len(self.specs)} (same-topology resume only)" - ) - for spec, meta in zip(self.specs, manifest["buckets"]): - assert meta["numel"] == spec.numel - assert meta["entry_numels"] == [main.numel() for _, main, _ in spec.entries] - shutil.copyfile(os.path.join(dirpath, meta["file"]), spec.path) - for group, step in zip(spec.adam.param_groups, meta["steps"]): - if step: - group["step"] = step - for _, main, _ in spec.entries: - state = spec.adam.state.setdefault(main, {}) - for key in _MOMENT_KEYS: - if key not in state: - t = torch.empty_like(main) - t.untyped_storage().resize_(0) - state[key] = t - spec.main_on_disk = True - spec.moments_on_disk = True - fp32_state_path = os.path.join(dirpath, "fp32_resident_optimizer.pt") - if self._fp32_adam is not None and os.path.isfile(fp32_state_path): - self._fp32_adam.load_state_dict(torch.load(fp32_state_path)) - logger.info(f"NVMe optimizer state loaded: {len(self.specs)} buckets <- {dirpath}") - - @torch.no_grad() - def refresh_main_from_model_params(self, copy_fn) -> None: - for spec in self.specs: - for tensor, _ in self._segment(spec, "main"): - self._materialize(tensor) - copy_fn() - for spec in self.specs: - for tensor, offset in self._segment(spec, "main"): - self._stream(spec.fd, offset, tensor, to_disk=True) - self._release(tensor) - spec.main_on_disk = True - - @torch.no_grad() - def step(self) -> bool: - t0 = time.monotonic() - read_bytes = written_bytes = 0 - for spec in self.specs: - read_bytes += self._load_bucket(spec) - self._sync_hyperparams(spec) - spec.adam.step() - self.dist_opt._copy_main_params_to_model_params_for( - (main, model) for model, main, _ in spec.entries - ) - written_bytes += self._store_bucket(spec) - if self._fp32_adam is not None: - self._sync_fp32_hyperparams() - self._fp32_adam.step() - logger.info( - f"NVMe streaming step: {len(self.specs)} buckets, " - f"read {read_bytes / 1024**3:.1f} GB, wrote {written_bytes / 1024**3:.1f} GB " - f"in {time.monotonic() - t0:.1f}s" - ) - return True - - def _sync_hyperparams(self, spec: "_BucketSpec") -> None: - master_groups = self.dist_opt.optimizer.param_groups - for group, g_idx in zip(spec.adam.param_groups, spec.group_master_indices): - group["lr"] = master_groups[g_idx]["lr"] - group["weight_decay"] = master_groups[g_idx]["weight_decay"] - - def _sync_fp32_hyperparams(self) -> None: - master_groups = self.dist_opt.optimizer.param_groups - for group, g_idx in zip(self._fp32_adam.param_groups, self._fp32_group_indices): - group["lr"] = master_groups[g_idx]["lr"] - group["weight_decay"] = master_groups[g_idx]["weight_decay"] - - def _load_bucket(self, spec: "_BucketSpec") -> int: - nbytes = 0 - if spec.main_on_disk: - for tensor, offset in self._segment(spec, "main"): - self._materialize(tensor) - self._stream(spec.fd, offset, tensor, to_disk=False) - nbytes += tensor.numel() * 4 - if spec.moments_on_disk: - for key in _MOMENT_KEYS: - for tensor, offset in self._segment(spec, key): - self._materialize(tensor) - self._stream(spec.fd, offset, tensor, to_disk=False) - nbytes += tensor.numel() * 4 - return nbytes - - def _store_bucket(self, spec: "_BucketSpec") -> int: - nbytes = 0 - for key in ("main",) + _MOMENT_KEYS: - for tensor, offset in self._segment(spec, key): - self._stream(spec.fd, offset, tensor, to_disk=True) - self._release(tensor) - nbytes += tensor.numel() * 4 - spec.main_on_disk = True - spec.moments_on_disk = True - return nbytes - - # ------------------------------------------------------- residency & I/O - - def _segment(self, spec: "_BucketSpec", key: str): - segment_index = ("main",) + _MOMENT_KEYS - base = segment_index.index(key) * spec.numel * 4 - for (_, main, _), entry_offset in zip(spec.entries, spec.entry_offsets): - tensor = main if key == "main" else spec.adam.state[main][key] - yield tensor, base + entry_offset * 4 - - @staticmethod - def _materialize(tensor: torch.Tensor) -> None: - tensor.untyped_storage().resize_(tensor.numel() * tensor.element_size()) - - @staticmethod - def _release(tensor: torch.Tensor) -> None: - tensor.untyped_storage().resize_(0) - - def _stream(self, fd: int, base_offset: int, tensor: torch.Tensor, *, to_disk: bool) -> None: - flat = tensor.view(-1) - chunk_numel = self._chunk.numel() - pos = 0 - while pos < flat.numel(): - n = min(chunk_numel, flat.numel() - pos) - byte_offset = base_offset + pos * 4 - if to_disk: - self._chunk[:n].copy_(flat[pos : pos + n]) - self._rw_full(os.pwritev, fd, byte_offset, self._chunk_np[:n]) - else: - self._rw_full(os.preadv, fd, byte_offset, self._chunk_np[:n]) - flat[pos : pos + n].copy_(self._chunk[:n]) - pos += n - - @staticmethod - def _rw_full(op, fd: int, offset: int, array) -> None: - mv = memoryview(array).cast("B") - done = 0 - while done < len(mv): - n = op(fd, [mv[done:]], offset + done) - if n <= 0: - raise IOError(f"short {op.__name__} ({n}) on optimizer state file at offset {offset + done}") - done += n From 34cd1bf2a4113f734c4538037995d75ad5c856ed Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Fri, 24 Jul 2026 12:02:14 -0700 Subject: [PATCH 08/12] Point the NVMe store import at miles The move commit deleted megatron/core/optimizer/nvme_state_store.py but left the import behind, so enabling --optimizer-state-nvme-dir raised ImportError. --- megatron/core/optimizer/distrib_optimizer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 77c1410f01d..985c8a04444 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -632,7 +632,7 @@ def __init__( self._state_offloader = OptimizerStateOffloader(self) if self.config.optimizer_state_nvme_dir is not None: - from megatron.core.optimizer.nvme_state_store import NVMeOptimizerStateStore + from miles.backends.megatron_utils.nvme_state_store import NVMeOptimizerStateStore self._nvme_state_store = NVMeOptimizerStateStore( self, From 1a2581a216750f11ee74b560389852462799f0ef Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Fri, 24 Jul 2026 12:23:19 -0700 Subject: [PATCH 09/12] Add --optimizer-state-nvme-moment-dtype The step is bound by streaming volume, and the moments tolerate less precision than the master copy: bf16 storage cuts 12 bytes per parameter to 8. Defaults to fp32, which is bit-identical to keeping them on GPU. --- megatron/core/optimizer/optimizer_config.py | 4 ++++ megatron/training/arguments.py | 6 ++++++ 2 files changed, 10 insertions(+) diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index a62a9a98173..20aa363748b 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -340,6 +340,10 @@ class OptimizerConfig: optimizer_state_nvme_chunk_mb: int = 256 """Pinned staging chunk size for NVMe optimizer state streaming.""" + optimizer_state_nvme_moment_dtype: str = "fp32" + """Storage dtype for the NVMe-streamed Adam moments. bf16 cuts streaming volume by a + third, which the step is bound by; fp32 is bit-identical to keeping them on GPU.""" + ################ # Miscellaneous ################ diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 9c7d2e3fa0b..60e73f3ff5b 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2117,6 +2117,12 @@ def _add_training_args(parser): 'state is not supported yet.') group.add_argument('--optimizer-state-nvme-chunk-mb', type=int, default=256, help='Pinned staging chunk size for NVMe optimizer state streaming.') + group.add_argument('--optimizer-state-nvme-moment-dtype', type=str, default='fp32', + choices=['fp32', 'bf16', 'fp16', 'fp8e4m3', 'fp8e5m2'], + help='Storage dtype for the NVMe-streamed Adam moments. bf16 cuts ' + 'streaming volume by a third, which the step is bound by; fp32 is ' + 'bit-identical to keeping them on GPU. The fp8 options need ' + 'per-block scaling to be sound and are not recommended.') group.add_argument('--dataloader-type', type=str, default=None, choices=['single', 'cyclic', 'external'], help='Single pass vs multiple pass data loader') From 86edf228289d56a515873147b6ec7cb5fc916549 Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Fri, 24 Jul 2026 14:59:53 -0700 Subject: [PATCH 10/12] Drop the docstring from _copy_main_params_to_model_params_for --- megatron/core/optimizer/distrib_optimizer.py | 1 - 1 file changed, 1 deletion(-) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 985c8a04444..db6b55dfe63 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -2509,7 +2509,6 @@ def copy_group_params(shard_main_groups, model_groups): copy_group_params(self.shard_fp32_groups, self.model_fp32_groups) def _copy_main_params_to_model_params_for(self, pairs): - """Copy (shard_main_param, model_param) pairs into the param buffer.""" for shard_main_param, model_param in pairs: param_range_map = self._get_model_param_range_map(model_param) world_range = param_range_map["gbuf_world_in_bucket"] From 3bf38312ebc9b6e4109420a50713eb9d06e81651 Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Fri, 24 Jul 2026 16:30:59 -0700 Subject: [PATCH 11/12] Keep the NVMe state store in Megatron The store only manipulates DistributedOptimizer internals -- bucket layout, param groups, Adam state, main-param residency -- so it belongs next to them rather than across a repo boundary. --- megatron/core/optimizer/distrib_optimizer.py | 2 +- megatron/core/optimizer/nvme_state_store.py | 359 +++++++++++++++++++ 2 files changed, 360 insertions(+), 1 deletion(-) create mode 100644 megatron/core/optimizer/nvme_state_store.py diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index db6b55dfe63..703a74c9cf3 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -632,7 +632,7 @@ def __init__( self._state_offloader = OptimizerStateOffloader(self) if self.config.optimizer_state_nvme_dir is not None: - from miles.backends.megatron_utils.nvme_state_store import NVMeOptimizerStateStore + from megatron.core.optimizer.nvme_state_store import NVMeOptimizerStateStore self._nvme_state_store = NVMeOptimizerStateStore( self, diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py new file mode 100644 index 00000000000..e859518d0a9 --- /dev/null +++ b/megatron/core/optimizer/nvme_state_store.py @@ -0,0 +1,359 @@ +import atexit +import errno +import json +import logging +import os +import shutil +import time +from typing import TYPE_CHECKING, NamedTuple + +import torch + +if TYPE_CHECKING: + from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer + +logger = logging.getLogger(__name__) + +SEGMENTS = ("main", "exp_avg", "exp_avg_sq") +DTYPES = { + "fp32": torch.float32, + "bf16": torch.bfloat16, + "fp16": torch.float16, + "fp8e4m3": torch.float8_e4m3fn, + "fp8e5m2": torch.float8_e5m2, +} +BUCKET_NUMEL_LIMIT = 200_000_000 +FP32_RESIDENT_WARN_MB = 256 +IO_ALIGN = 4096 + + +class _Entry(NamedTuple): + model_param: torch.nn.Parameter + main_param: torch.Tensor + group_index: int + + +def _align(nbytes: int) -> int: + return (nbytes + IO_ALIGN - 1) // IO_ALIGN * IO_ALIGN + + +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) + 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) + return fd + + +def _rw_full(op, fd: int, offset: int, buf) -> None: + mv = memoryview(buf).cast("B") + done = 0 + while done < len(mv): + n = op(fd, [mv[done:]], offset + done) + if n <= 0: + raise OSError(f"short {op.__name__} ({n}) on optimizer state file at offset {offset + done}") + done += n + + +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]): + for entry in entries: + current.append(entry) + numel += entry.main_param.numel() + if numel >= limit: + planned.append(current) + current, numel = [], 0 + if current: + planned.append(current) + current, numel = [], 0 + return planned + + +class _Stager: + def __init__(self, nbytes: int): + self._buf = torch.empty(nbytes, dtype=torch.uint8, pin_memory=True) + self._bytes = self._buf.numpy() + self._device_buf = None + + def _device_staging(self, dtype: torch.dtype, numel: int, like: torch.Tensor) -> torch.Tensor: + # A cross-dtype copy between GPU and pinned host memory does not take the DMA + # path: casting on the device first and moving same-dtype bytes is ~40x faster + # for bf16 and ~100x for fp8. + size = self._buf.numel() + if self._device_buf is None or self._device_buf.device != like.device: + self._device_buf = torch.empty(size, dtype=torch.uint8, device=like.device) + return self._device_buf[: numel * dtype.itemsize].view(dtype) + + def transfer(self, fd: int, offset: int, tensor: torch.Tensor, dtype: torch.dtype, *, to_disk: bool) -> int: + flat = tensor.view(-1) + cast = dtype != flat.dtype + chunk = self._buf.numel() // dtype.itemsize + pos = 0 + while pos < flat.numel(): + numel = min(chunk, flat.numel() - pos) + host = self._buf[: numel * dtype.itemsize].view(dtype) + at = offset + pos * dtype.itemsize + nbytes = numel * dtype.itemsize + if to_disk: + if cast: + staged = self._device_staging(dtype, numel, flat) + staged.copy_(flat[pos : pos + numel]) + host.copy_(staged) + else: + host.copy_(flat[pos : pos + numel]) + _rw_full(os.pwritev, fd, at, self._bytes[:nbytes]) + else: + _rw_full(os.preadv, fd, at, self._bytes[:nbytes]) + if cast: + staged = self._device_staging(dtype, numel, flat) + staged.copy_(host) + flat[pos : pos + numel].copy_(staged) + else: + flat[pos : pos + numel].copy_(host) + pos += numel + return flat.numel() * dtype.itemsize + + +class _Bucket: + def __init__(self, path: str, entries: list[_Entry], adam, stager: _Stager, dtypes: dict): + self.path, self.entries, self.adam, self.dtypes = path, entries, adam, dtypes + self._stager = stager + self.group_indices = sorted({e.group_index for e in entries}) + self.numel = sum(e.main_param.numel() for e in entries) + + self.offsets: dict[str, list[int]] = {} + at = 0 + for segment in SEGMENTS: + self.offsets[segment] = [] + for entry in entries: + self.offsets[segment].append(at) + at += _align(entry.main_param.numel() * dtypes[segment].itemsize) + self.nbytes = at + self.fd = _allocate_file(path, at) + self.moments_ready = False + + def _tensors(self, segment: str): + for index, entry in enumerate(self.entries): + tensor = entry.main_param if segment == "main" else self.adam.state[entry.main_param][segment] + yield tensor, self.offsets[segment][index] + + def _move(self, segments, *, to_disk: bool) -> int: + moved = 0 + for segment in segments: + for tensor, offset in self._tensors(segment): + if not to_disk: + _resize(tensor, tensor.numel()) + moved += self._stager.transfer(self.fd, offset, tensor, self.dtypes[segment], to_disk=to_disk) + if to_disk: + _resize(tensor, 0) + return moved + + def fetch(self) -> int: + return self._move(SEGMENTS if self.moments_ready else SEGMENTS[:1], to_disk=False) + + def flush(self, segments=SEGMENTS) -> int: + moved = self._move(segments, to_disk=True) + self.moments_ready = self.moments_ready or tuple(segments) == SEGMENTS + return moved + + def materialize_main(self) -> None: + for tensor, _ in self._tensors("main"): + _resize(tensor, tensor.numel()) + + def allocate_moments(self) -> None: + for entry in self.entries: + state = self.adam.state.setdefault(entry.main_param, {}) + for segment in SEGMENTS[1:]: + if segment not in state: + state[segment] = torch.empty_like(entry.main_param) + _resize(state[segment], 0) + self.moments_ready = True + + +class NVMeOptimizerStateStore: + _next_uid = 0 + + def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int): + self.dist_opt = distrib_optimizer + self.uid = NVMeOptimizerStateStore._next_uid + NVMeOptimizerStateStore._next_uid += 1 + config = distrib_optimizer.config + + assert not config.use_precision_aware_optimizer, ( + "NVMe state store requires the non-precision-aware optimizer " "(fp32 main params held by mcore)." + ) + assert not config.optimizer_cpu_offload, "NVMe state store is mutually exclusive with CPU offload." + assert ( + not config.offload_optimizer_states + ), "NVMe state store is mutually exclusive with --offload-optimizer-states." + assert not distrib_optimizer.ddp_config.use_megatron_fsdp + + moments = getattr(config, "optimizer_state_nvme_moment_dtype", "fp32") + assert moments in DTYPES, f"unknown moment dtype {moments!r}, expected one of {sorted(DTYPES)}" + if DTYPES[moments].itemsize == 1: + logger.warning( + f"Storing Adam moments as {moments} without per-block scaling is numerically risky " + "for exp_avg_sq; bf16 is the safe way to halve moment I/O." + ) + self.dtypes = {"main": torch.float32, "exp_avg": DTYPES[moments], "exp_avg_sq": DTYPES[moments]} + + rank = torch.distributed.get_rank() + instance = distrib_optimizer.distributed_optimizer_instance_id + self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}_{self.uid}") + shutil.rmtree(self.dir, ignore_errors=True) + os.makedirs(self.dir, exist_ok=True) + atexit.register(shutil.rmtree, self.dir, ignore_errors=True) + + self._stager = _Stager(chunk_mb * 1024 * 1024) + self.buckets = self._build_buckets() + self._fp32_group_indices, self._fp32_adam = self._build_fp32_optimizer() + + for bucket in self.buckets: + bucket.flush(segments=("main",)) + + total_gb = sum(b.nbytes for b in self.buckets) / 1024**3 + logger.info( + f"NVMe optimizer state store: {len(self.buckets)} buckets, {total_gb:.1f} GB at " + f"{self.dir} (moments stored as {moments})" + ) + + def _build_buckets(self) -> list[_Bucket]: + by_ddp_bucket: dict[tuple, list[_Entry]] = {} + groups = zip(self.dist_opt.model_float16_groups, self.dist_opt.shard_fp32_from_float16_groups, strict=True) + for group_index, (model_group, main_group) in enumerate(groups): + for model_param, main_param in zip(model_group, main_group, strict=True): + assert main_param is not None and main_param.dtype == torch.float32 + key = self.dist_opt.model_param_gbuf_map[model_param] + by_ddp_bucket.setdefault(key, []).append(_Entry(model_param, main_param, group_index)) + + buckets = [] + for index, entries in enumerate(plan_buckets(by_ddp_bucket)): + params: dict[int, list[torch.Tensor]] = {} + for entry in entries: + params.setdefault(entry.group_index, []).append(entry.main_param) + path = os.path.join(self.dir, f"bucket{index:05d}.bin") + buckets.append(_Bucket(path, entries, self._adam_for(params), self._stager, self.dtypes)) + return buckets + + def _build_fp32_optimizer(self): + params: dict[int, list[torch.Tensor]] = {} + total_bytes = 0 + for group_index, (model_group, shard_group) in enumerate( + zip(self.dist_opt.model_fp32_groups, self.dist_opt.shard_fp32_groups, strict=True) + ): + if model_group: + params[group_index] = list(shard_group) + total_bytes += sum(p.numel() * p.element_size() for p in shard_group) + if not params: + return [], None + + total_mb = total_bytes / 1024**2 + log = logger.warning if total_mb > FP32_RESIDENT_WARN_MB else logger.info + log(f"NVMe optimizer state store: {total_mb:.1f} MB of native-fp32 params stay GPU-resident") + return sorted(params), self._adam_for(params) + + def _adam_for(self, params_by_group: dict[int, list[torch.Tensor]]): + from megatron.core.optimizer import Adam + + master_groups = self.dist_opt.optimizer.param_groups + groups = [] + for group_index in sorted(params_by_group): + group = {k: v for k, v in master_groups[group_index].items() if k != "params"} + group["params"] = params_by_group[group_index] + groups.append(group) + return Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) + + def _sync_lr_wd(self, adam, group_indices) -> None: + master_groups = self.dist_opt.optimizer.param_groups + for group, group_index in zip(adam.param_groups, group_indices, strict=True): + group["lr"] = master_groups[group_index]["lr"] + group["weight_decay"] = master_groups[group_index]["weight_decay"] + + @torch.no_grad() + def step(self) -> bool: + started = time.monotonic() + read = written = 0 + for bucket in self.buckets: + read += bucket.fetch() + self._sync_lr_wd(bucket.adam, bucket.group_indices) + bucket.adam.step() + self.dist_opt._copy_main_params_to_model_params_for( + (entry.main_param, entry.model_param) for entry in bucket.entries + ) + written += bucket.flush() + if self._fp32_adam is not None: + self._sync_lr_wd(self._fp32_adam, self._fp32_group_indices) + self._fp32_adam.step() + logger.info( + f"NVMe streaming step: {len(self.buckets)} buckets, read {read / 1024**3:.1f} GB, " + f"wrote {written / 1024**3:.1f} GB in {time.monotonic() - started:.1f}s" + ) + return True + + @torch.no_grad() + def refresh_main_from_model_params(self, copy_fn) -> None: + for bucket in self.buckets: + bucket.materialize_main() + copy_fn() + for bucket in self.buckets: + bucket.flush(segments=("main",)) + + @torch.no_grad() + def save_to(self, dirpath: str) -> None: + os.makedirs(dirpath, exist_ok=True) + manifest = { + "dtypes": {segment: str(dtype) for segment, dtype in self.dtypes.items()}, + "buckets": [ + { + "numel": bucket.numel, + "entry_numels": [e.main_param.numel() for e in bucket.entries], + "steps": [g.get("step", 0) for g in bucket.adam.param_groups], + "file": os.path.basename(bucket.path), + } + for bucket in self.buckets + ], + } + for bucket in self.buckets: + shutil.copyfile(bucket.path, os.path.join(dirpath, os.path.basename(bucket.path))) + if self._fp32_adam is not None: + torch.save(self._fp32_adam.state_dict(), os.path.join(dirpath, "fp32_resident_optimizer.pt")) + with open(os.path.join(dirpath, "manifest.json"), "w") as f: + json.dump(manifest, f) + logger.info(f"NVMe optimizer state saved: {len(self.buckets)} buckets -> {dirpath}") + + @torch.no_grad() + def load_from(self, dirpath: str) -> None: + with open(os.path.join(dirpath, "manifest.json")) as f: + manifest = json.load(f) + + saved = manifest.get("dtypes", {segment: str(torch.float32) for segment in SEGMENTS}) + current = {segment: str(dtype) for segment, dtype in self.dtypes.items()} + assert saved == current, ( + f"NVMe state dtype mismatch: checkpoint stores {saved}, this run stores {current} " + "-- the bytes would be misread" + ) + assert len(manifest["buckets"]) == len(self.buckets), ( + f"NVMe state layout mismatch: checkpoint has {len(manifest['buckets'])} buckets, " + f"current topology builds {len(self.buckets)} (same-topology resume only)" + ) + + for bucket, meta in zip(self.buckets, manifest["buckets"], strict=True): + assert meta["numel"] == bucket.numel + assert meta["entry_numels"] == [e.main_param.numel() for e in bucket.entries] + shutil.copyfile(os.path.join(dirpath, meta["file"]), bucket.path) + for group, step in zip(bucket.adam.param_groups, meta["steps"], strict=True): + if step: + group["step"] = step + bucket.allocate_moments() + fp32_state = os.path.join(dirpath, "fp32_resident_optimizer.pt") + if self._fp32_adam is not None and os.path.isfile(fp32_state): + self._fp32_adam.load_state_dict(torch.load(fp32_state)) + logger.info(f"NVMe optimizer state loaded: {len(self.buckets)} buckets <- {dirpath}") From 058116646940b9ccae0b6d6c70ca6636641ea3bc Mon Sep 17 00:00:00 2001 From: yueming-yuan Date: Fri, 24 Jul 2026 17:33:08 -0700 Subject: [PATCH 12/12] Move the NVMe state store to miles, keep only the two checkpoint hooks The store and its wiring move to miles_plugins/optimizers/nvme_stream.py (radixark/miles#1793), which binds the five entry points that used to be `if self._nvme_state_store is not None` branches in distrib_optimizer.py onto each DistributedOptimizer instance. The --optimizer-state-nvme-* args and their OptimizerConfig fields go with it; miles owns the flags now and passes them to the constructor. distrib_optimizer.py is back to untouched. The earlier _copy_main_params_to_model_params_for extraction is reverted too: its only live caller under streaming was the store, and it handed out a partial operation, since the fp8 branch depends on quantize_param_shard() having run over the whole fp8 param set -- a DP collective that cannot be split per bucket. miles adapts the copy instead, with the source pinned by commit and line range, and rejects fp8 params up front. What is left here is the two checkpoint hooks, duck-typed through the _nvme_state_store attribute and importing nothing from miles. They have to be inside save_checkpoint/load_checkpoint to cover every call path: miles has two save call sites and a missed one shows up as a checkpoint with no optimizer state that silently resumes from zero. save_to/load_from now take the checkpoint base and derive their own per-rank subdirectory from the same layout the live scratch directory uses, so _nvme_state_checkpoint_dir goes away and nothing here reads store.dist_opt or store.uid. The duck-typed contract is four methods. The "this checkpoint has no streamed state" case moves into load_from, where the policy belongs. --- megatron/core/optimizer/distrib_optimizer.py | 63 +--- megatron/core/optimizer/nvme_state_store.py | 359 ------------------- megatron/core/optimizer/optimizer_config.py | 11 - megatron/training/arguments.py | 13 - megatron/training/checkpointing.py | 16 +- 5 files changed, 22 insertions(+), 440 deletions(-) delete mode 100644 megatron/core/optimizer/nvme_state_store.py diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 703a74c9cf3..8fe58f92bbb 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -537,7 +537,6 @@ def __init__( ) self._state_offloader: Optional[OptimizerStateOffloader] = None - self._nvme_state_store = None # when freezing sub-models we have no real optimizer # but still need a stub DistributedOptimizer class @@ -631,15 +630,6 @@ def __init__( if self.config.offload_optimizer_states: self._state_offloader = OptimizerStateOffloader(self) - if self.config.optimizer_state_nvme_dir is not None: - from megatron.core.optimizer.nvme_state_store import NVMeOptimizerStateStore - - self._nvme_state_store = NVMeOptimizerStateStore( - self, - self.config.optimizer_state_nvme_dir, - self.config.optimizer_state_nvme_chunk_mb, - ) - def _get_model_param_range_map(self, param: torch.nn.Parameter): """ Given a model param, get the index sub-range of the param that this @@ -666,8 +656,6 @@ def state_dict(self): optimizer state (e.g., exp_avg, exp_avg_sq) are stored in a separate checkpoint file by calling 'save_parameter_state()'. """ - if self._nvme_state_store is not None: - return {"nvme_state_store": True} inner_state_dict = self.optimizer.state_dict() state_dict = {} @@ -753,9 +741,6 @@ def load_state_dict(self, state_dict): - state_order : The index of a parameter within the shared parameter list. """ - if self._nvme_state_store is not None: - return - if self.ddp_config.use_megatron_fsdp: if "param_to_group_meta" in state_dict: state_dict["param_groups"] = self._param2group_meta_to_param_groups( @@ -1262,8 +1247,6 @@ def sharded_state_dict( Regular state dict parameters are saved on DP rank 0 and loaded on all ranks. """ - if self._nvme_state_store is not None: - return {} if sharding_type is not None: log_single_rank( logger, @@ -2503,27 +2486,29 @@ def _copy_main_params_to_model_params(self): # Utility method for copying group params. def copy_group_params(shard_main_groups, model_groups): for shard_main_group, model_group in zip(shard_main_groups, model_groups): - self._copy_main_params_to_model_params_for(zip(shard_main_group, model_group)) + for shard_main_param, model_param in zip(shard_main_group, model_group): - copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups) - copy_group_params(self.shard_fp32_groups, self.model_fp32_groups) + param_range_map = self._get_model_param_range_map(model_param) + world_range = param_range_map["gbuf_world_in_bucket"] - def _copy_main_params_to_model_params_for(self, pairs): - for shard_main_param, model_param in pairs: - param_range_map = self._get_model_param_range_map(model_param) - world_range = param_range_map["gbuf_world_in_bucket"] + assert world_range.size == shard_main_param.nelement() - assert world_range.size == shard_main_param.nelement() + gbuf_index, _, bucket_id = self.model_param_gbuf_map[model_param] + model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data - gbuf_index, _, bucket_id = self.model_param_gbuf_map[model_param] - model_param_buffer = self.buffers[gbuf_index].buckets[bucket_id].param_data + shard_model_param = model_param_buffer.view(-1)[ + world_range.start : world_range.end + ] - shard_model_param = model_param_buffer.view(-1)[world_range.start : world_range.end] + if is_float8tensor(model_param): + # FP8 params are quantized in the above "quantize_param_shard" function. + continue + else: + shard_model_param.data.copy_(shard_main_param) - if is_float8tensor(model_param): - # FP8 params are quantized in the above "quantize_param_shard" function. - continue - shard_model_param.data.copy_(shard_main_param) + # Copy shard groups to model groups. + copy_group_params(self.shard_fp32_from_float16_groups, self.model_float16_groups) + copy_group_params(self.shard_fp32_groups, self.model_fp32_groups) def _copy_main_params_to_param_buffer(self): """ @@ -2586,15 +2571,6 @@ def _build_model_param_to_state_dict_param_map(self, state_dict): return model_param_to_state_dict_param_map def _copy_model_params_to_main_params(self, state_dict=None): - if self._nvme_state_store is not None: - self._nvme_state_store.refresh_main_from_model_params( - lambda: self._copy_model_params_to_main_params_impl(state_dict) - ) - return - self._copy_model_params_to_main_params_impl(state_dict) - - @torch.no_grad() - def _copy_model_params_to_main_params_impl(self, state_dict=None): """ Copy model params to main params. @@ -2658,10 +2634,7 @@ def step_with_ready_grads(self) -> bool: """ if self._state_offloader is not None: self._state_offloader.sync_before_step() - if self._nvme_state_store is not None: - update_successful = self._nvme_state_store.step() - else: - update_successful = super().step_with_ready_grads() + update_successful = super().step_with_ready_grads() timers = self.config.timers if timers is not None: diff --git a/megatron/core/optimizer/nvme_state_store.py b/megatron/core/optimizer/nvme_state_store.py deleted file mode 100644 index e859518d0a9..00000000000 --- a/megatron/core/optimizer/nvme_state_store.py +++ /dev/null @@ -1,359 +0,0 @@ -import atexit -import errno -import json -import logging -import os -import shutil -import time -from typing import TYPE_CHECKING, NamedTuple - -import torch - -if TYPE_CHECKING: - from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer - -logger = logging.getLogger(__name__) - -SEGMENTS = ("main", "exp_avg", "exp_avg_sq") -DTYPES = { - "fp32": torch.float32, - "bf16": torch.bfloat16, - "fp16": torch.float16, - "fp8e4m3": torch.float8_e4m3fn, - "fp8e5m2": torch.float8_e5m2, -} -BUCKET_NUMEL_LIMIT = 200_000_000 -FP32_RESIDENT_WARN_MB = 256 -IO_ALIGN = 4096 - - -class _Entry(NamedTuple): - model_param: torch.nn.Parameter - main_param: torch.Tensor - group_index: int - - -def _align(nbytes: int) -> int: - return (nbytes + IO_ALIGN - 1) // IO_ALIGN * IO_ALIGN - - -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) - 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) - return fd - - -def _rw_full(op, fd: int, offset: int, buf) -> None: - mv = memoryview(buf).cast("B") - done = 0 - while done < len(mv): - n = op(fd, [mv[done:]], offset + done) - if n <= 0: - raise OSError(f"short {op.__name__} ({n}) on optimizer state file at offset {offset + done}") - done += n - - -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]): - for entry in entries: - current.append(entry) - numel += entry.main_param.numel() - if numel >= limit: - planned.append(current) - current, numel = [], 0 - if current: - planned.append(current) - current, numel = [], 0 - return planned - - -class _Stager: - def __init__(self, nbytes: int): - self._buf = torch.empty(nbytes, dtype=torch.uint8, pin_memory=True) - self._bytes = self._buf.numpy() - self._device_buf = None - - def _device_staging(self, dtype: torch.dtype, numel: int, like: torch.Tensor) -> torch.Tensor: - # A cross-dtype copy between GPU and pinned host memory does not take the DMA - # path: casting on the device first and moving same-dtype bytes is ~40x faster - # for bf16 and ~100x for fp8. - size = self._buf.numel() - if self._device_buf is None or self._device_buf.device != like.device: - self._device_buf = torch.empty(size, dtype=torch.uint8, device=like.device) - return self._device_buf[: numel * dtype.itemsize].view(dtype) - - def transfer(self, fd: int, offset: int, tensor: torch.Tensor, dtype: torch.dtype, *, to_disk: bool) -> int: - flat = tensor.view(-1) - cast = dtype != flat.dtype - chunk = self._buf.numel() // dtype.itemsize - pos = 0 - while pos < flat.numel(): - numel = min(chunk, flat.numel() - pos) - host = self._buf[: numel * dtype.itemsize].view(dtype) - at = offset + pos * dtype.itemsize - nbytes = numel * dtype.itemsize - if to_disk: - if cast: - staged = self._device_staging(dtype, numel, flat) - staged.copy_(flat[pos : pos + numel]) - host.copy_(staged) - else: - host.copy_(flat[pos : pos + numel]) - _rw_full(os.pwritev, fd, at, self._bytes[:nbytes]) - else: - _rw_full(os.preadv, fd, at, self._bytes[:nbytes]) - if cast: - staged = self._device_staging(dtype, numel, flat) - staged.copy_(host) - flat[pos : pos + numel].copy_(staged) - else: - flat[pos : pos + numel].copy_(host) - pos += numel - return flat.numel() * dtype.itemsize - - -class _Bucket: - def __init__(self, path: str, entries: list[_Entry], adam, stager: _Stager, dtypes: dict): - self.path, self.entries, self.adam, self.dtypes = path, entries, adam, dtypes - self._stager = stager - self.group_indices = sorted({e.group_index for e in entries}) - self.numel = sum(e.main_param.numel() for e in entries) - - self.offsets: dict[str, list[int]] = {} - at = 0 - for segment in SEGMENTS: - self.offsets[segment] = [] - for entry in entries: - self.offsets[segment].append(at) - at += _align(entry.main_param.numel() * dtypes[segment].itemsize) - self.nbytes = at - self.fd = _allocate_file(path, at) - self.moments_ready = False - - def _tensors(self, segment: str): - for index, entry in enumerate(self.entries): - tensor = entry.main_param if segment == "main" else self.adam.state[entry.main_param][segment] - yield tensor, self.offsets[segment][index] - - def _move(self, segments, *, to_disk: bool) -> int: - moved = 0 - for segment in segments: - for tensor, offset in self._tensors(segment): - if not to_disk: - _resize(tensor, tensor.numel()) - moved += self._stager.transfer(self.fd, offset, tensor, self.dtypes[segment], to_disk=to_disk) - if to_disk: - _resize(tensor, 0) - return moved - - def fetch(self) -> int: - return self._move(SEGMENTS if self.moments_ready else SEGMENTS[:1], to_disk=False) - - def flush(self, segments=SEGMENTS) -> int: - moved = self._move(segments, to_disk=True) - self.moments_ready = self.moments_ready or tuple(segments) == SEGMENTS - return moved - - def materialize_main(self) -> None: - for tensor, _ in self._tensors("main"): - _resize(tensor, tensor.numel()) - - def allocate_moments(self) -> None: - for entry in self.entries: - state = self.adam.state.setdefault(entry.main_param, {}) - for segment in SEGMENTS[1:]: - if segment not in state: - state[segment] = torch.empty_like(entry.main_param) - _resize(state[segment], 0) - self.moments_ready = True - - -class NVMeOptimizerStateStore: - _next_uid = 0 - - def __init__(self, distrib_optimizer: "DistributedOptimizer", dir_root: str, chunk_mb: int): - self.dist_opt = distrib_optimizer - self.uid = NVMeOptimizerStateStore._next_uid - NVMeOptimizerStateStore._next_uid += 1 - config = distrib_optimizer.config - - assert not config.use_precision_aware_optimizer, ( - "NVMe state store requires the non-precision-aware optimizer " "(fp32 main params held by mcore)." - ) - assert not config.optimizer_cpu_offload, "NVMe state store is mutually exclusive with CPU offload." - assert ( - not config.offload_optimizer_states - ), "NVMe state store is mutually exclusive with --offload-optimizer-states." - assert not distrib_optimizer.ddp_config.use_megatron_fsdp - - moments = getattr(config, "optimizer_state_nvme_moment_dtype", "fp32") - assert moments in DTYPES, f"unknown moment dtype {moments!r}, expected one of {sorted(DTYPES)}" - if DTYPES[moments].itemsize == 1: - logger.warning( - f"Storing Adam moments as {moments} without per-block scaling is numerically risky " - "for exp_avg_sq; bf16 is the safe way to halve moment I/O." - ) - self.dtypes = {"main": torch.float32, "exp_avg": DTYPES[moments], "exp_avg_sq": DTYPES[moments]} - - rank = torch.distributed.get_rank() - instance = distrib_optimizer.distributed_optimizer_instance_id - self.dir = os.path.join(dir_root, f"rank{rank}", f"opt{instance}_{self.uid}") - shutil.rmtree(self.dir, ignore_errors=True) - os.makedirs(self.dir, exist_ok=True) - atexit.register(shutil.rmtree, self.dir, ignore_errors=True) - - self._stager = _Stager(chunk_mb * 1024 * 1024) - self.buckets = self._build_buckets() - self._fp32_group_indices, self._fp32_adam = self._build_fp32_optimizer() - - for bucket in self.buckets: - bucket.flush(segments=("main",)) - - total_gb = sum(b.nbytes for b in self.buckets) / 1024**3 - logger.info( - f"NVMe optimizer state store: {len(self.buckets)} buckets, {total_gb:.1f} GB at " - f"{self.dir} (moments stored as {moments})" - ) - - def _build_buckets(self) -> list[_Bucket]: - by_ddp_bucket: dict[tuple, list[_Entry]] = {} - groups = zip(self.dist_opt.model_float16_groups, self.dist_opt.shard_fp32_from_float16_groups, strict=True) - for group_index, (model_group, main_group) in enumerate(groups): - for model_param, main_param in zip(model_group, main_group, strict=True): - assert main_param is not None and main_param.dtype == torch.float32 - key = self.dist_opt.model_param_gbuf_map[model_param] - by_ddp_bucket.setdefault(key, []).append(_Entry(model_param, main_param, group_index)) - - buckets = [] - for index, entries in enumerate(plan_buckets(by_ddp_bucket)): - params: dict[int, list[torch.Tensor]] = {} - for entry in entries: - params.setdefault(entry.group_index, []).append(entry.main_param) - path = os.path.join(self.dir, f"bucket{index:05d}.bin") - buckets.append(_Bucket(path, entries, self._adam_for(params), self._stager, self.dtypes)) - return buckets - - def _build_fp32_optimizer(self): - params: dict[int, list[torch.Tensor]] = {} - total_bytes = 0 - for group_index, (model_group, shard_group) in enumerate( - zip(self.dist_opt.model_fp32_groups, self.dist_opt.shard_fp32_groups, strict=True) - ): - if model_group: - params[group_index] = list(shard_group) - total_bytes += sum(p.numel() * p.element_size() for p in shard_group) - if not params: - return [], None - - total_mb = total_bytes / 1024**2 - log = logger.warning if total_mb > FP32_RESIDENT_WARN_MB else logger.info - log(f"NVMe optimizer state store: {total_mb:.1f} MB of native-fp32 params stay GPU-resident") - return sorted(params), self._adam_for(params) - - def _adam_for(self, params_by_group: dict[int, list[torch.Tensor]]): - from megatron.core.optimizer import Adam - - master_groups = self.dist_opt.optimizer.param_groups - groups = [] - for group_index in sorted(params_by_group): - group = {k: v for k, v in master_groups[group_index].items() if k != "params"} - group["params"] = params_by_group[group_index] - groups.append(group) - return Adam(groups, adam_w_mode=self.dist_opt.config.decoupled_weight_decay) - - def _sync_lr_wd(self, adam, group_indices) -> None: - master_groups = self.dist_opt.optimizer.param_groups - for group, group_index in zip(adam.param_groups, group_indices, strict=True): - group["lr"] = master_groups[group_index]["lr"] - group["weight_decay"] = master_groups[group_index]["weight_decay"] - - @torch.no_grad() - def step(self) -> bool: - started = time.monotonic() - read = written = 0 - for bucket in self.buckets: - read += bucket.fetch() - self._sync_lr_wd(bucket.adam, bucket.group_indices) - bucket.adam.step() - self.dist_opt._copy_main_params_to_model_params_for( - (entry.main_param, entry.model_param) for entry in bucket.entries - ) - written += bucket.flush() - if self._fp32_adam is not None: - self._sync_lr_wd(self._fp32_adam, self._fp32_group_indices) - self._fp32_adam.step() - logger.info( - f"NVMe streaming step: {len(self.buckets)} buckets, read {read / 1024**3:.1f} GB, " - f"wrote {written / 1024**3:.1f} GB in {time.monotonic() - started:.1f}s" - ) - return True - - @torch.no_grad() - def refresh_main_from_model_params(self, copy_fn) -> None: - for bucket in self.buckets: - bucket.materialize_main() - copy_fn() - for bucket in self.buckets: - bucket.flush(segments=("main",)) - - @torch.no_grad() - def save_to(self, dirpath: str) -> None: - os.makedirs(dirpath, exist_ok=True) - manifest = { - "dtypes": {segment: str(dtype) for segment, dtype in self.dtypes.items()}, - "buckets": [ - { - "numel": bucket.numel, - "entry_numels": [e.main_param.numel() for e in bucket.entries], - "steps": [g.get("step", 0) for g in bucket.adam.param_groups], - "file": os.path.basename(bucket.path), - } - for bucket in self.buckets - ], - } - for bucket in self.buckets: - shutil.copyfile(bucket.path, os.path.join(dirpath, os.path.basename(bucket.path))) - if self._fp32_adam is not None: - torch.save(self._fp32_adam.state_dict(), os.path.join(dirpath, "fp32_resident_optimizer.pt")) - with open(os.path.join(dirpath, "manifest.json"), "w") as f: - json.dump(manifest, f) - logger.info(f"NVMe optimizer state saved: {len(self.buckets)} buckets -> {dirpath}") - - @torch.no_grad() - def load_from(self, dirpath: str) -> None: - with open(os.path.join(dirpath, "manifest.json")) as f: - manifest = json.load(f) - - saved = manifest.get("dtypes", {segment: str(torch.float32) for segment in SEGMENTS}) - current = {segment: str(dtype) for segment, dtype in self.dtypes.items()} - assert saved == current, ( - f"NVMe state dtype mismatch: checkpoint stores {saved}, this run stores {current} " - "-- the bytes would be misread" - ) - assert len(manifest["buckets"]) == len(self.buckets), ( - f"NVMe state layout mismatch: checkpoint has {len(manifest['buckets'])} buckets, " - f"current topology builds {len(self.buckets)} (same-topology resume only)" - ) - - for bucket, meta in zip(self.buckets, manifest["buckets"], strict=True): - assert meta["numel"] == bucket.numel - assert meta["entry_numels"] == [e.main_param.numel() for e in bucket.entries] - shutil.copyfile(os.path.join(dirpath, meta["file"]), bucket.path) - for group, step in zip(bucket.adam.param_groups, meta["steps"], strict=True): - if step: - group["step"] = step - bucket.allocate_moments() - fp32_state = os.path.join(dirpath, "fp32_resident_optimizer.pt") - if self._fp32_adam is not None and os.path.isfile(fp32_state): - self._fp32_adam.load_state_dict(torch.load(fp32_state)) - logger.info(f"NVMe optimizer state loaded: {len(self.buckets)} buckets <- {dirpath}") diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 20aa363748b..0f7081f4fc1 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -333,17 +333,6 @@ class OptimizerConfig: low_memory_resume: bool = False """If True, allocate optimizer states on CPU during checkpoint loading to prevent GPU OOM.""" - optimizer_state_nvme_dir: Optional[str] = None - """If set, stream fp32 main params and Adam moments through per-bucket files under this - node-local directory during the optimizer step, bounding GPU residency to one bucket.""" - - optimizer_state_nvme_chunk_mb: int = 256 - """Pinned staging chunk size for NVMe optimizer state streaming.""" - - optimizer_state_nvme_moment_dtype: str = "fp32" - """Storage dtype for the NVMe-streamed Adam moments. bf16 cuts streaming volume by a - third, which the step is bound by; fp32 is bit-identical to keeping them on GPU.""" - ################ # Miscellaneous ################ diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 60e73f3ff5b..28fea46195d 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2110,19 +2110,6 @@ def _add_training_args(parser): 'Only support TE FusedAdam optimizer.' 'Note that this still uses pure GPU optimizer instead of ' 'HybridDeviceOptimizer for --optimizer-cpu-offload.') - group.add_argument('--optimizer-state-nvme-dir', type=str, default=None, - help='Stream fp32 main params and Adam moments through per-bucket ' - 'files under this node-local directory during the optimizer step, ' - 'bounding GPU residency to one bucket. Checkpointing optimizer ' - 'state is not supported yet.') - group.add_argument('--optimizer-state-nvme-chunk-mb', type=int, default=256, - help='Pinned staging chunk size for NVMe optimizer state streaming.') - group.add_argument('--optimizer-state-nvme-moment-dtype', type=str, default='fp32', - choices=['fp32', 'bf16', 'fp16', 'fp8e4m3', 'fp8e5m2'], - help='Storage dtype for the NVMe-streamed Adam moments. bf16 cuts ' - 'streaming volume by a third, which the step is bound by; fp32 is ' - 'bit-identical to keeping them on GPU. The fp8 options need ' - 'per-block scaling to be sound and are not recommended.') group.add_argument('--dataloader-type', type=str, default=None, choices=['single', 'cyclic', 'external'], help='Single pass vs multiple pass data loader') diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 0ea6101d0cb..6b47211d8d2 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -475,12 +475,8 @@ def _iter_nvme_state_stores(optimizer): yield store -def _nvme_state_checkpoint_dir(checkpoint_name, store): - rank = torch.distributed.get_rank() - instance = store.dist_opt.distributed_optimizer_instance_id - return os.path.join( - checkpoint_name, "nvme_opt_state", f"rank{rank:04d}_opt{instance}_{store.uid}" - ) +def _nvme_state_base_dir(checkpoint_name): + return os.path.join(checkpoint_name, "nvme_opt_state") def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floating_point_operations_so_far, @@ -587,7 +583,7 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati if not args.no_save_optim and optimizer is not None: for store in _iter_nvme_state_stores(optimizer): - store.save_to(_nvme_state_checkpoint_dir(checkpoint_name, store)) + store.save_to(_nvme_state_base_dir(checkpoint_name)) async_save_request = None if args.async_save: @@ -1850,11 +1846,7 @@ def load_model_state_dict(module, state_dict, strict: bool): if optimizer is not None and not release and not args.finetune and not args.no_load_optim: for store in _iter_nvme_state_stores(optimizer): - nvme_dir = _nvme_state_checkpoint_dir(checkpoint_name, store) - if os.path.isdir(nvme_dir): - store.load_from(nvme_dir) - else: - print_rank_0(f" no NVMe optimizer state at {nvme_dir}; starting fresh") + store.load_from(_nvme_state_base_dir(checkpoint_name)) # rerun state if not ignore_rerun_state: