From b6537e4c0fb6cc2d86c0795f45424071ef0d0d32 Mon Sep 17 00:00:00 2001 From: Zhichenzzz Date: Tue, 14 Jul 2026 20:56:08 -0700 Subject: [PATCH] [tml] inkling model support --- .../dist_checkpointing/strategies/torch.py | 17 +- megatron/core/optimizer/distrib_optimizer.py | 65 ++-- megatron/core/optimizer/nvme_state_store.py | 291 ++++++++++++++++++ megatron/core/optimizer/optimizer_config.py | 10 + megatron/core/transformer/moe/experts.py | 35 +++ .../core/transformer/moe/token_dispatcher.py | 6 + .../core/transformer/transformer_config.py | 25 +- megatron/training/arguments.py | 7 + megatron/training/checkpointing.py | 32 ++ 9 files changed, 468 insertions(+), 20 deletions(-) create mode 100644 megatron/core/optimizer/nvme_state_store.py diff --git a/megatron/core/dist_checkpointing/strategies/torch.py b/megatron/core/dist_checkpointing/strategies/torch.py index fb0bcaa4b3c..a74194ccdc3 100644 --- a/megatron/core/dist_checkpointing/strategies/torch.py +++ b/megatron/core/dist_checkpointing/strategies/torch.py @@ -109,6 +109,20 @@ def register_default_torch_strategies(): logger = getLogger(__name__) +_COORDINATION_PROCESS_GROUP = None + + +def _get_coordination_process_group(): + """Gloo group for checkpoint plan/metadata object collectives. On an + NCCL-only default group these would be staged through GPU and leave + permanent per-peer NCCL channel buffers on the coordinator rank.""" + global _COORDINATION_PROCESS_GROUP + if not torch.distributed.is_initialized(): + return None + if _COORDINATION_PROCESS_GROUP is None: + _COORDINATION_PROCESS_GROUP = torch.distributed.new_group(backend="gloo") + return _COORDINATION_PROCESS_GROUP + def flatten_state_dict( state_dict: ShardedStateDict, @@ -701,7 +715,7 @@ def async_save( ) = save_state_dict_async_plan( pyt_state_dict, writer, - None, + _get_coordination_process_group(), coordinator, planner=MCoreSavePlanner( dedup_replicated_tensors=not self.keep_only_main_replica, flatten_state_dict=False @@ -805,6 +819,7 @@ def load(self, sharded_state_dict: ShardedStateDict, checkpoint_dir: Path) -> St checkpoint.load_state_dict( pyt_state_dict, fsr, + process_group=_get_coordination_process_group(), planner=MCoreLoadPlanner( shapes_validation_sharded_tensors=flexible_shape_sharded_tensors, allow_shape_mismatch_sharded_tensors=allow_shape_mismatch_sharded_tensors, diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 8fe58f92bbb..2c45f178d5d 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,8 @@ 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 = {} @@ -741,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( @@ -1247,6 +1262,8 @@ 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, @@ -2486,29 +2503,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): """ @@ -2571,6 +2588,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. @@ -2634,7 +2660,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..e21acb8c961 --- /dev/null +++ b/megatron/core/optimizer/nvme_state_store.py @@ -0,0 +1,291 @@ +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") + + +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 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, ( + "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}_{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() + 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) + + # ------------------------------------------------------------------ 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: + 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) + 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/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index d8e75342226..d3ab5301595 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -526,6 +526,33 @@ def backward_dw(self): pass +class _MoEActivationInFP32(torch.autograd.Function): + """fp32 swiglu (x prob) between the grouped GEMMs, one round back to the params dtype.""" + + @staticmethod + def forward(ctx, fc1_out, probs, glu_offset): + ctx.save_for_backward(fc1_out, probs) + ctx.glu_offset = glu_offset + g, u = torch.chunk(fc1_out.float(), 2, dim=-1) + y = F.silu(g) * (u + glu_offset) * probs.float() + return y.to(fc1_out.dtype) + + @staticmethod + def backward(ctx, grad_out): + fc1_out, probs = ctx.saved_tensors + g, u = torch.chunk(fc1_out.float(), 2, dim=-1) + s = torch.sigmoid(g) + silu = g * s + go = grad_out.float() * probs.float() + d_g = go * (u + ctx.glu_offset) * (s + silu * (1 - s)) + d_u = go * silu + d_p = None + if ctx.needs_input_grad[1]: + d_p = (grad_out.float() * silu * (u + ctx.glu_offset)).sum(-1, keepdim=True) + d_p = d_p.to(probs.dtype) + return torch.cat([d_g, d_u], dim=-1).to(fc1_out.dtype), d_p, None + + class TEGroupedMLP(MegatronModule): """An efficient implementation of the Experts layer using TE's GroupedLinear. @@ -681,6 +708,9 @@ def forward( # Probs already applied, so reset to 1. permuted_probs = torch.ones_like(permuted_probs) + if self.config.moe_combine_in_fp32: + permuted_probs = torch.ones_like(permuted_probs) + with off_interface( self.offload_expert_fc1, permuted_local_hidden_states, "expert_fc1" ) as permuted_local_hidden_states: @@ -695,6 +725,11 @@ def forward( ) def bias_act_func(intermediate_parallel, bias_parallel, permuted_probs): + if self.config.moe_activation_in_fp32: + assert bias_parallel is None and self.config.gated_linear_unit + return _MoEActivationInFP32.apply( + intermediate_parallel, permuted_probs, self.config.glu_linear_offset + ) if self.config.use_te_activation_func: if bias_parallel is not None: intermediate_parallel = intermediate_parallel + bias_parallel diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 327dbc8a382..e3b65e91709 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -832,11 +832,17 @@ def combine_postprocess(self, permutated_local_input_tokens): self.shared_experts.post_forward_comm() # Unpermutation 1: AlltoAll output to output + if self.config.moe_combine_in_fp32: + assert not self.config.moe_permute_fusion + combine_probs = self.probs + else: + combine_probs = None output = unpermute( permutated_local_input_tokens, self.reversed_local_input_permutation_mapping, restore_shape=self.hidden_shape_before_permute, routing_map=self.routing_map, + probs=combine_probs, fused=self.config.moe_permute_fusion, drop_and_pad=self.drop_and_pad, ) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 5c495ea8b2c..4ffb527a98a 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -691,6 +691,12 @@ class TransformerConfig(ModelParallelConfig): improve stability especially when the number of experts is large (e.g. finegrained-moe). None means no changes for dtype.""" + moe_activation_in_fp32: bool = False + """Compute the inter-GEMM swiglu in fp32 with one round back to the params dtype.""" + + moe_combine_in_fp32: bool = False + """Apply routing probs and accumulate expert outputs in fp32 with one final round.""" + moe_router_enable_expert_bias: bool = False """TopK routing with dynamic per-expert bias in the aux-loss-free load balancing strategy. The routing decision is based on the sum of the routing scores and the expert bias. @@ -983,7 +989,9 @@ def __post_init__(self): """ super().__post_init__() if self.true_on_policy_contract is not None: - from miles_megatron_plugins.true_on_policy.contracts import validate_true_on_policy_contract + from miles_megatron_plugins.true_on_policy.contracts import ( + validate_true_on_policy_contract, + ) validate_true_on_policy_contract(self.true_on_policy_contract) if self.fp16 and self.bf16: @@ -1922,6 +1930,21 @@ def __post_init__(self): f"variable sequence length, please use alltoall dispatcher instead." ) + if self.moe_activation_in_fp32 or self.moe_combine_in_fp32: + assert not ( + self.fp8 or getattr(self, 'fp4', None) + ), "moe_*_in_fp32 supports bf16/fp16 only" + assert ( + not self.moe_permute_fusion + ), "moe_*_in_fp32 bypasses fused permute/unpermute; disable --moe-permute-fusion" + assert ( + self.activation_func_clamp_value is None + ), "moe_*_in_fp32 is an inference-alignment mode; drop --activation-func-clamp-value" + if self.moe_combine_in_fp32: + assert ( + self.moe_router_dtype == 'fp32' + ), "moe_combine_in_fp32 needs --moe-router-dtype fp32 (probs reach combine un-rounded)" + if self.moe_permute_fusion: from megatron.core.transformer.moe.moe_utils import ( fused_permute, 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') diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 7c87eca191a..a6464fa50df 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -468,6 +468,21 @@ 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}_{store.uid}" + ) + + 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 +585,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 +1850,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: