diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 8fe58f92bbb..77c1410f01d 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,28 @@ 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_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 +2587,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 +2659,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/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 0f7081f4fc1..a62a9a98173 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -333,6 +333,13 @@ 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.""" + ################ # 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') diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 7c87eca191a..0ea6101d0cb 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,10 @@ 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) + 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,14 @@ def load_model_state_dict(module, state_dict, strict: bool): else: optimizer.reload_model_params() + 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: