diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 1441a71518d..04a93e9f376 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -1247,6 +1247,7 @@ def _load_global_dist_base_checkpoint( sharded_state_dict, checkpoint_name, load_strategy, + validate_access_integrity=args.ckpt_load_validate_sharding_integrity, strict=args.dist_ckpt_strictness, ) return state_dict, checkpoint_name, release, CheckpointType.GLOBAL diff --git a/megatron/training/config/training_config.py b/megatron/training/config/training_config.py index 27cffb837f4..b44a6472872 100644 --- a/megatron/training/config/training_config.py +++ b/megatron/training/config/training_config.py @@ -529,6 +529,11 @@ class CheckpointConfig: ckpt_assume_constant_structure: bool = False """Assume the checkpoint structure is constant across saves to enable optimizations.""" + ckpt_load_validate_sharding_integrity: bool = True + """Whether to validate sharding access integrity when loading a distributed checkpoint. + When True (default), each tensor shard is checked to be accessed exactly once as main + replica by some rank. Disabling skips this validation""" + strict_fsdp_dtensor_load: bool = True """Whether to enforce strict loading for FSDP DTensor checkpoints. When False, allows partial loading.""" 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 5f9c617893c..f8aace0105d 100644 --- a/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py +++ b/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py @@ -140,6 +140,7 @@ def create_args(): args.vocab_file = None args.add_position_embedding = False args.ckpt_assume_constant_structure = True + args.ckpt_load_validate_sharding_integrity = True args.dist_ckpt_strictness = "assume_ok_unexpected" args.fp16 = False args.bf16 = True diff --git a/tests/unit_tests/dist_checkpointing/utils.py b/tests/unit_tests/dist_checkpointing/utils.py index 0aadaee3b29..8a9df54ddc8 100644 --- a/tests/unit_tests/dist_checkpointing/utils.py +++ b/tests/unit_tests/dist_checkpointing/utils.py @@ -150,6 +150,7 @@ def init_checkpointing_mock_args(args, ckpt_dir, fully_parallel=False): args.no_save_optim = False args.no_save_rng = False args.ckpt_assume_constant_structure = False + args.ckpt_load_validate_sharding_integrity = True args.log_progress = False args.auto_detect_ckpt_format = False args.exit_on_missing_checkpoint = False diff --git a/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py b/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py index a6afabe8817..f871938b218 100644 --- a/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py +++ b/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py @@ -135,6 +135,7 @@ def create_args(): args.vocab_file = None args.add_position_embedding = False args.ckpt_assume_constant_structure = False + args.ckpt_load_validate_sharding_integrity = True args.dist_ckpt_strictness = "assume_ok_unexpected" args.fp16 = False args.bf16 = True diff --git a/tests/unit_tests/test_checkpointing.py b/tests/unit_tests/test_checkpointing.py index 16bac10566d..61cda3d91d2 100644 --- a/tests/unit_tests/test_checkpointing.py +++ b/tests/unit_tests/test_checkpointing.py @@ -140,6 +140,7 @@ def create_ckpt_load_args(create_args): args.ckpt_assume_constant_structure = False args.ckpt_fully_parallel_save = False args.ckpt_fully_parallel_load = False + args.ckpt_load_validate_sharding_integrity = True args.dist_ckpt_strictness = 'assume_ok_unexpected' args.use_megatron_fsdp = False args.strict_fsdp_dtensor_load = True