@@ -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
4973def select_rng_state (rng_state , rng_key : str , data_parallel_random_init : bool ):
0 commit comments