diff --git a/megatron/core/MSC_Integration.md b/megatron/core/MSC_Integration.md index da8b5c982b8..cd44d6afbb2 100644 --- a/megatron/core/MSC_Integration.md +++ b/megatron/core/MSC_Integration.md @@ -125,14 +125,16 @@ python pretrain_gpt.py \ **Notes:** Only the `torch_dist` checkpoint format is currently supported when saving to or loading from MSC URLs. -## Disable MSC +## Enable MSC -By default, MSC integration is automatically enabled when the `multi-storage-client` library is installed. MSC is also used for regular filesystem paths (like `/filesystem_mountpoint/path` in `--data-path`, `--save`, or `--load`) even when not using explicit MSC URLs. MSC functions as a very thin abstraction layer with negligible performance impact when used with regular paths, so there's typically no need to disable it. If you need to disable MSC, you can do so using the `--disable-msc` flag: +MSC integration is opt-in: even when the `multi-storage-client` library is installed, MSC is **disabled by default**. To opt in, pass the `--enable-msc` flag. Once enabled, MSC is also used for regular filesystem paths (like `/filesystem_mountpoint/path` in `--data-path`, `--save`, or `--load`), not just explicit `msc://` URLs. ```bash -python pretrain_gpt.py --disable-msc +python pretrain_gpt.py --enable-msc ``` +> **Note:** When MSC is enabled, the dist-checkpointing loader uses `msc.torch.MultiStorageFileSystemReader` instead of `CachedMetadataFileSystemReader`. This means `ckpt_assume_constant_structure=True` (and any other path that requests `cache_metadata=True`) will be silently overridden — metadata is re-read on every load. A warning is emitted in this case. + ## Performance Considerations When using object storage with MSC, there are a few important performance implications to keep in mind: diff --git a/megatron/core/dist_checkpointing/strategies/torch.py b/megatron/core/dist_checkpointing/strategies/torch.py index 7943561700f..31782acb851 100644 --- a/megatron/core/dist_checkpointing/strategies/torch.py +++ b/megatron/core/dist_checkpointing/strategies/torch.py @@ -831,6 +831,14 @@ def _get_filesystem_reader( ) -> FileSystemReader: if MultiStorageClientFeature.is_enabled(): msc = MultiStorageClientFeature.import_package() + if cache_metadata: + warnings.warn( + "MSC is enabled: returning msc.torch.MultiStorageFileSystemReader instead of " + "CachedMetadataFileSystemReader. The cache_metadata=True request " + "(e.g. ckpt_assume_constant_structure=True) will be ignored and metadata " + "will be re-read on every load. Pass --enable-msc only when this is intended.", + stacklevel=2, + ) return msc.torch.MultiStorageFileSystemReader(checkpoint_dir, thread_count=2) if cache_metadata: diff --git a/megatron/core/msc_utils.py b/megatron/core/msc_utils.py index b8d85a0dfd9..ce7cb685e25 100644 --- a/megatron/core/msc_utils.py +++ b/megatron/core/msc_utils.py @@ -8,11 +8,9 @@ try: import multistorageclient as msc - _msc_available = True logger.info('The multistorageclient package is available.') except ModuleNotFoundError: msc = None - _msc_available = False class _FeatureFlag: @@ -41,7 +39,7 @@ def import_package(self) -> Any: ) if not self.is_enabled(): raise RuntimeError( - "The MSC feature is disabled. Please enable by removing the --disable-msc argument." + "The MSC feature is disabled. Please enable it by passing --enable-msc." ) return msc @@ -54,7 +52,7 @@ def __setstate__(self, state): self._enabled = state['_enabled'] -MultiStorageClientFeature = _FeatureFlag(_msc_available) +MultiStorageClientFeature = _FeatureFlag(default=False) def open_file(*args, **kwargs): diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 9f68bec074b..6fe84b7b3cf 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -147,11 +147,27 @@ def parse_args(extra_args_provider=None, ignore_unknown_args=False): args.rank = int(os.getenv('RANK', '0')) args.world_size = int(os.getenv("WORLD_SIZE", '1')) - # Args to disable MSC - if not args.enable_msc: + # Args to enable MSC (opt-in: disabled by default) + if args.disable_msc_deprecated: + warn_rank_0( + '--disable-msc is deprecated and will be removed in a future release. ' + 'MSC is now disabled by default; pass --enable-msc to opt in.' + ) + # Preserve legacy semantics: --disable-msc forces MSC off, even if + # --enable-msc was also passed. + args.enable_msc = False + + if args.enable_msc: + MultiStorageClientFeature.enable() + if not MultiStorageClientFeature.is_enabled(): + raise RuntimeError( + "--enable-msc was passed but the multistorageclient package is not " + "installed. Install it with `pip install multi-storage-client`." + ) + warn_rank_0('The MSC feature is enabled.') + else: MultiStorageClientFeature.disable() assert MultiStorageClientFeature.is_enabled() is False - warn_rank_0('The MSC feature is disabled.') return args @@ -3195,8 +3211,13 @@ def _add_experimental_args(parser): def _add_msc_args(parser): group = parser.add_argument_group(title="msc") - group.add_argument('--disable-msc', default=True, action='store_false', dest='enable_msc', - help='Disable the usage of Multi-Storage Client (MSC) in Megatron Core.') + group.add_argument('--enable-msc', default=False, action='store_true', dest='enable_msc', + help='Enable the usage of Multi-Storage Client (MSC) in Megatron Core. ' + 'Disabled by default; pass this flag to opt in.') + group.add_argument('--disable-msc', default=False, action='store_true', + dest='disable_msc_deprecated', + help='[DEPRECATED] MSC is disabled by default; this flag is a no-op ' + 'and will be removed in a future release.') return parser def _add_kitchen_quantization_arguments(parser: argparse.ArgumentParser): diff --git a/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py b/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py index 3e493539b1b..4ed91aa2cb6 100644 --- a/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py +++ b/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py @@ -7,6 +7,9 @@ import torch from megatron.core import mpu +from megatron.core.dist_checkpointing.strategies.cached_metadata_filesystem_reader import ( + CachedMetadataFileSystemReader, +) from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_with_transformer_engine_spec as gpt_te_spec, @@ -293,6 +296,10 @@ def test_save_and_load_checkpoint_vpp( args.async_strategy = "mcore" set_args(args) + # Cached metadata is mutated in-place by torch's load planner, so reusing it + # across loads of the same checkpoint dir corrupts subsequent reads. + CachedMetadataFileSystemReader.clear_metadata_cache() + def set_tp_pp_vpp(tp, pp, vpp=None, pp_layout=None, destroy_first=True): if destroy_first: Utils.destroy_model_parallel() @@ -306,45 +313,16 @@ def set_ckpt_path(ckpt_path): args.save = ckpt_path args.load = ckpt_path - set_tp_pp_vpp(*src_tp_pp_vpp, pp_layout=src_pp_layout, destroy_first=False) - init_num_microbatches_calculator( - rank=0, global_batch_size=1, micro_batch_size=1, data_parallel_size=1 - ) - - iteration = 123 - layer_spec_fn = get_gpt_decoder_block_spec if is_moe else gpt_te_spec - model = initialize_gpt_model( - 1, - layer_spec_fn=layer_spec_fn, - num_layers=args.num_layers, - hidden_size=args.hidden_size, - num_attention_heads=args.num_attention_heads, - tensor_model_parallel_size=args.tensor_model_parallel_size, - pipeline_model_parallel_size=args.pipeline_model_parallel_size, - virtual_pipeline_model_parallel_size=args.virtual_pipeline_model_parallel_size, - pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, - is_moe=is_moe, - ) - model = model if isinstance(model, list) else [model] - optimizer = None - opt_param_scheduler = None - num_floating_point_operations_so_far = 456 - - with ( - TempNamedDir(tmp_path_dist_ckpt / 'test_gpt_model_reconfiguration_model_A') as ckpt_dir_A, - TempNamedDir(tmp_path_dist_ckpt / 'test_gpt_model_reconfiguration_model_B') as ckpt_dir_B, - ): - set_ckpt_path(ckpt_dir_A) - save_checkpoint( - iteration, model, optimizer, opt_param_scheduler, num_floating_point_operations_so_far + try: + set_tp_pp_vpp(*src_tp_pp_vpp, pp_layout=src_pp_layout, destroy_first=False) + init_num_microbatches_calculator( + rank=0, global_batch_size=1, micro_batch_size=1, data_parallel_size=1 ) - expected_ckpt_path = args.save / "iter_0000123" / ".metadata" - assert os.path.exists(expected_ckpt_path) - - set_tp_pp_vpp(*dst_tp_pp_vpp, pp_layout=dst_pp_layout) - new_model = initialize_gpt_model( - 2, + iteration = 123 + layer_spec_fn = get_gpt_decoder_block_spec if is_moe else gpt_te_spec + model = initialize_gpt_model( + 1, layer_spec_fn=layer_spec_fn, num_layers=args.num_layers, hidden_size=args.hidden_size, @@ -355,57 +333,95 @@ def set_ckpt_path(ckpt_path): pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, is_moe=is_moe, ) - new_model = new_model if isinstance(new_model, list) else [new_model] + model = model if isinstance(model, list) else [model] + optimizer = None + opt_param_scheduler = None + num_floating_point_operations_so_far = 456 - load_checkpoint(new_model, optimizer, opt_param_scheduler, strict=False) - set_ckpt_path(ckpt_dir_B) - save_checkpoint( - iteration, - new_model, - optimizer, - opt_param_scheduler, - num_floating_point_operations_so_far, - ) + with ( + TempNamedDir( + tmp_path_dist_ckpt / 'test_gpt_model_reconfiguration_model_A' + ) as ckpt_dir_A, + TempNamedDir( + tmp_path_dist_ckpt / 'test_gpt_model_reconfiguration_model_B' + ) as ckpt_dir_B, + ): + set_ckpt_path(ckpt_dir_A) + save_checkpoint( + iteration, + model, + optimizer, + opt_param_scheduler, + num_floating_point_operations_so_far, + ) - set_tp_pp_vpp(1, 1) - set_ckpt_path(ckpt_dir_A) - model_A = initialize_gpt_model( - 123, - layer_spec_fn=layer_spec_fn, - num_layers=args.num_layers, - hidden_size=args.hidden_size, - num_attention_heads=args.num_attention_heads, - tensor_model_parallel_size=args.tensor_model_parallel_size, - pipeline_model_parallel_size=args.pipeline_model_parallel_size, - virtual_pipeline_model_parallel_size=args.virtual_pipeline_model_parallel_size, - pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, - is_moe=is_moe, - ) - load_checkpoint([model_A], optimizer, opt_param_scheduler, strict=False) + expected_ckpt_path = args.save / "iter_0000123" / ".metadata" + assert os.path.exists(expected_ckpt_path) - set_ckpt_path(ckpt_dir_B) - model_B = initialize_gpt_model( - 321, - layer_spec_fn=layer_spec_fn, - num_layers=args.num_layers, - hidden_size=args.hidden_size, - num_attention_heads=args.num_attention_heads, - tensor_model_parallel_size=args.tensor_model_parallel_size, - pipeline_model_parallel_size=args.pipeline_model_parallel_size, - virtual_pipeline_model_parallel_size=args.virtual_pipeline_model_parallel_size, - pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, - is_moe=is_moe, - ) - load_checkpoint([model_B], optimizer, opt_param_scheduler, strict=False) + set_tp_pp_vpp(*dst_tp_pp_vpp, pp_layout=dst_pp_layout) + new_model = initialize_gpt_model( + 2, + layer_spec_fn=layer_spec_fn, + num_layers=args.num_layers, + hidden_size=args.hidden_size, + num_attention_heads=args.num_attention_heads, + tensor_model_parallel_size=args.tensor_model_parallel_size, + pipeline_model_parallel_size=args.pipeline_model_parallel_size, + virtual_pipeline_model_parallel_size=args.virtual_pipeline_model_parallel_size, + pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, + is_moe=is_moe, + ) + new_model = new_model if isinstance(new_model, list) else [new_model] + + load_checkpoint(new_model, optimizer, opt_param_scheduler, strict=False) + set_ckpt_path(ckpt_dir_B) + save_checkpoint( + iteration, + new_model, + optimizer, + opt_param_scheduler, + num_floating_point_operations_so_far, + ) + + set_tp_pp_vpp(1, 1) + set_ckpt_path(ckpt_dir_A) + model_A = initialize_gpt_model( + 123, + layer_spec_fn=layer_spec_fn, + num_layers=args.num_layers, + hidden_size=args.hidden_size, + num_attention_heads=args.num_attention_heads, + tensor_model_parallel_size=args.tensor_model_parallel_size, + pipeline_model_parallel_size=args.pipeline_model_parallel_size, + virtual_pipeline_model_parallel_size=args.virtual_pipeline_model_parallel_size, + pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, + is_moe=is_moe, + ) + load_checkpoint([model_A], optimizer, opt_param_scheduler, strict=False) - for k in model_A.state_dict(): - if "_extra_state" in k: # Ignore extra states - continue - tensor_a = model_A.state_dict()[k] - tensor_b = model_B.state_dict()[k] - assert tensor_a is not None, k - assert tensor_b is not None, k - assert torch.equal(tensor_a, tensor_b), k + set_ckpt_path(ckpt_dir_B) + model_B = initialize_gpt_model( + 321, + layer_spec_fn=layer_spec_fn, + num_layers=args.num_layers, + hidden_size=args.hidden_size, + num_attention_heads=args.num_attention_heads, + tensor_model_parallel_size=args.tensor_model_parallel_size, + pipeline_model_parallel_size=args.pipeline_model_parallel_size, + virtual_pipeline_model_parallel_size=args.virtual_pipeline_model_parallel_size, + pipeline_model_parallel_layout=args.pipeline_model_parallel_layout, + is_moe=is_moe, + ) + load_checkpoint([model_B], optimizer, opt_param_scheduler, strict=False) - Utils.destroy_model_parallel() - unset_num_microbatches_calculator() + for k in model_A.state_dict(): + if "_extra_state" in k: # Ignore extra states + continue + tensor_a = model_A.state_dict()[k] + tensor_b = model_B.state_dict()[k] + assert tensor_a is not None, k + assert tensor_b is not None, k + assert torch.equal(tensor_a, tensor_b), k + finally: + Utils.destroy_model_parallel() + unset_num_microbatches_calculator()