Skip to content

Commit 7785517

Browse files
committed
Fix FSDP RNG restore across DP resharding
1 parent eeef573 commit 7785517

2 files changed

Lines changed: 51 additions & 14 deletions

File tree

swift/megatron/utils/megatron_fsdp_checkpoint.py

Lines changed: 31 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -32,18 +32,42 @@ def build_rng_state(rng_state, data_parallel_random_init: bool = False, rng_key:
3232
return {rng_key: rng_state_list}
3333

3434

35-
def get_rng_load_key(checkpoint_dir: str) -> str:
35+
def get_rng_load_key(checkpoint_dir: str, data_parallel_random_init: bool = False) -> Optional[str]:
3636
metadata_keys = FileSystemReader(checkpoint_dir).read_metadata().state_dict_metadata
37-
rank_key = f'global_rank_{torch.distributed.get_rank()}'
38-
if any(key.startswith(f'rng_state.{rank_key}.') for key in metadata_keys):
39-
return rank_key
37+
global_rank_prefix = 'rng_state.global_rank_'
38+
saved_global_ranks = set()
39+
for key in metadata_keys:
40+
if key.startswith(global_rank_prefix):
41+
rank = key[len(global_rank_prefix):].split('.', 1)[0]
42+
if rank.isdigit():
43+
saved_global_ranks.add(int(rank))
44+
45+
if saved_global_ranks:
46+
# Per-rank RNG is exact only when the saved and current world-rank sets match.
47+
# A topology-changing reshard should still load model/optimizer state without RNG.
48+
current_world_size = torch.distributed.get_world_size()
49+
if saved_global_ranks != set(range(current_world_size)):
50+
return None
51+
current_rank = torch.distributed.get_rank()
52+
rank_key = f'global_rank_{current_rank}'
53+
return rank_key if current_rank in saved_global_ranks else None
4054

4155
pp_rank = mpu.get_pipeline_model_parallel_rank()
4256
tp_rank = mpu.get_tensor_model_parallel_rank()
4357
legacy_key = f'({pp_rank}, {tp_rank})'
44-
if any(key.startswith(f'rng_state.{legacy_key}.') for key in metadata_keys):
45-
return legacy_key
46-
raise RuntimeError(f'RNG state for global rank {torch.distributed.get_rank()} was not found in `{checkpoint_dir}`.')
58+
legacy_prefix = f'rng_state.{legacy_key}.'
59+
legacy_metadata_keys = [key for key in metadata_keys if key.startswith(legacy_prefix)]
60+
if not legacy_metadata_keys:
61+
return None
62+
if data_parallel_random_init:
63+
saved_dp_ranks = set()
64+
for key in legacy_metadata_keys:
65+
dp_rank = key[len(legacy_prefix):].split('.', 1)[0]
66+
if dp_rank.isdigit():
67+
saved_dp_ranks.add(int(dp_rank))
68+
if saved_dp_ranks != set(range(mpu.get_data_parallel_world_size())):
69+
return None
70+
return legacy_key
4771

4872

4973
def select_rng_state(rng_state, rng_key: str, data_parallel_random_init: bool):

swift/megatron/utils/megatron_lm_utils.py

Lines changed: 20 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -468,12 +468,25 @@ def load_mcore_checkpoint(args,
468468
if (ckpt_tp_pp == run_tp_pp and not finetune and not no_load_rng
469469
and not getattr(state_dict['args'], 'no_save_rng', False)):
470470
if fsdp_dtensor:
471-
fsdp_rng_key = fsdp_checkpoint.get_rng_load_key(checkpoint_dir)
472-
gen_sd_rng_state = _get_rng_state(
473-
fsdp_dtensor=fsdp_dtensor,
474-
data_parallel_random_init=args.data_parallel_random_init,
475-
fsdp_rng_key=fsdp_rng_key,
476-
) # we can load the rng state
471+
fsdp_rng_key = fsdp_checkpoint.get_rng_load_key(checkpoint_dir, args.data_parallel_random_init)
472+
if fsdp_rng_key is None:
473+
gen_sd_rng_state = None
474+
logger.warning(
475+
f'Megatron-FSDP RNG state in `{checkpoint_dir}` is incompatible with the current distributed '
476+
f'topology/world size. Model, optimizer, and scheduler checkpoint loading will continue, but '
477+
f'exact RNG restore is skipped; resumed training is not guaranteed to be bitwise or stepwise '
478+
f'deterministic.')
479+
else:
480+
gen_sd_rng_state = _get_rng_state(
481+
fsdp_dtensor=True,
482+
data_parallel_random_init=args.data_parallel_random_init,
483+
fsdp_rng_key=fsdp_rng_key,
484+
)
485+
else:
486+
gen_sd_rng_state = _get_rng_state(
487+
fsdp_dtensor=False,
488+
data_parallel_random_init=args.data_parallel_random_init,
489+
) # we can load the rng state
477490
else:
478491
gen_sd_rng_state = None
479492
if ckpt_tp_pp != run_tp_pp:
@@ -543,7 +556,7 @@ def load_mcore_checkpoint(args,
543556
elif (args.fp16 or args.bf16) and optimizer is not None:
544557
optimizer.reload_model_params()
545558

546-
if not finetune and not no_load_rng:
559+
if not finetune and not no_load_rng and (not fsdp_dtensor or fsdp_rng_key is not None):
547560
if 'rng_state' in state_dict:
548561
if fsdp_dtensor:
549562
rng_state = fsdp_checkpoint.select_rng_state(

0 commit comments

Comments
 (0)