Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions megatron/core/MSC_Integration.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
8 changes: 8 additions & 0 deletions megatron/core/dist_checkpointing/strategies/torch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
6 changes: 2 additions & 4 deletions megatron/core/msc_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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

Expand All @@ -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):
Expand Down
31 changes: 26 additions & 5 deletions megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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`."
)
Comment thread
asolergi-nv marked this conversation as resolved.
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

Expand Down Expand Up @@ -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',
Comment thread
asolergi-nv marked this conversation as resolved.
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):
Expand Down
188 changes: 102 additions & 86 deletions tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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()
Expand All @@ -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,
Expand All @@ -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()
Loading