Skip to content

Commit 26ced08

Browse files
Preserve Megatron RNG state across resume
1 parent f106742 commit 26ced08

6 files changed

Lines changed: 27 additions & 42 deletions

File tree

swift/megatron/trainers/base.py

Lines changed: 16 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -577,8 +577,12 @@ def copy_path(src_path: str, tgt_path: str):
577577
else:
578578
raise ValueError(f'Source path is neither a file nor a directory: {src_path}')
579579

580-
def _prepare_data_iterator(self, train_dataset, val_dataset=None, use_origin_cyclic: bool = False):
581-
train_dataloader, val_dataloader = self._prepare_dataloader(train_dataset, val_dataset)
580+
def _prepare_data_iterator(self,
581+
train_dataset,
582+
val_dataset=None,
583+
use_origin_cyclic: bool = False,
584+
seed: Optional[int] = None):
585+
train_dataloader, val_dataloader = self._prepare_dataloader(train_dataset, val_dataset, seed=seed)
582586
train_data_iterator = iter(self.cyclic_iter(train_dataloader, use_origin_cyclic=use_origin_cyclic))
583587
val_data_iterator = None
584588
if val_dataset is not None:
@@ -973,11 +977,15 @@ def _aggregated_metrics(self, metrics, total_metrics):
973977
total_metrics[key] = torch.tensor([0.0, 0.0], dtype=torch.float32, device=torch.cuda.current_device())
974978
total_metrics[key] += val
975979

976-
def _prepare_dataloader(self, train_dataset, val_dataset=None):
980+
def _prepare_dataloader(self, train_dataset, val_dataset=None, seed: Optional[int] = None):
977981
args = self.args
978982
val_dataloader = None
983+
generator = None
984+
if seed is not None:
985+
generator = torch.Generator()
986+
generator.manual_seed(seed)
979987
if args.streaming:
980-
train_dataloader = build_streaming_dataloader(args, train_dataset, self.data_collator)
988+
train_dataloader = build_streaming_dataloader(args, train_dataset, self.data_collator, generator=generator)
981989
if val_dataset is not None:
982990
val_dataloader = build_streaming_dataloader(args, val_dataset, self.data_collator)
983991
return train_dataloader, val_dataloader
@@ -991,8 +999,9 @@ def _prepare_dataloader(self, train_dataset, val_dataset=None):
991999
data_sharding=args.data_sharding,
9921000
shuffle=args.train_dataloader_shuffle,
9931001
group_by_length=args.group_by_length,
1002+
seed=seed or 0,
9941003
)
995-
train_dataloader = self._create_dataloader(train_dataset, train_batch_sampler)
1004+
train_dataloader = self._create_dataloader(train_dataset, train_batch_sampler, generator=generator)
9961005
if val_dataset is not None:
9971006
val_batch_sampler = MegatronPretrainingSampler(
9981007
total_samples=len(val_dataset),
@@ -1004,7 +1013,7 @@ def _prepare_dataloader(self, train_dataset, val_dataset=None):
10041013
val_dataloader = self._create_dataloader(val_dataset, val_batch_sampler)
10051014
return train_dataloader, val_dataloader
10061015

1007-
def _create_dataloader(self, dataset, batch_sampler):
1016+
def _create_dataloader(self, dataset, batch_sampler, generator=None):
10081017
args = self.args
10091018

10101019
dataloader = torch.utils.data.DataLoader(
@@ -1015,6 +1024,7 @@ def _create_dataloader(self, dataset, batch_sampler):
10151024
persistent_workers=args.dataloader_persistent_workers if args.dataloader_num_workers > 0 else False,
10161025
prefetch_factor=args.dataloader_prefetch_factor if args.dataloader_num_workers > 0 else None,
10171026
collate_fn=self.data_collator,
1027+
generator=generator,
10181028
)
10191029
return dataloader
10201030

swift/megatron/trainers/batch_sampler.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,7 @@ def __init__(
7575
data_sharding,
7676
shuffle: bool = True,
7777
group_by_length: bool = False,
78+
seed: int = 0,
7879
):
7980
# Keep a copy of input params for later use.
8081
self.dataset = dataset
@@ -93,6 +94,7 @@ def __init__(
9394
self.data_sharding = data_sharding
9495
self.shuffle = shuffle
9596
self.group_by_length = group_by_length
97+
self.seed = seed
9698
self.lengths = self.dataset['lengths'] if group_by_length else None
9799
if self.lengths is not None:
98100
self.lengths = [max(length) if isinstance(length, list) else length for length in self.lengths]
@@ -124,14 +126,14 @@ def __iter__(self):
124126
start_idx = self.data_parallel_rank * bucket_size
125127

126128
g = torch.Generator()
127-
g.manual_seed(self.epoch)
129+
g.manual_seed(self.seed + self.epoch)
128130
random_idx = torch.randperm(bucket_size, generator=g).tolist()
129131
idx_range = [start_idx + x for x in random_idx[bucket_offset:]]
130132
else:
131133
full_bucket_size = (self.total_samples // self.micro_batch_size) * self.micro_batch_size
132134
full_bucket_offset = current_epoch_samples
133135
g = torch.Generator()
134-
g.manual_seed(self.epoch)
136+
g.manual_seed(self.seed + self.epoch)
135137
if self.group_by_length:
136138
from transformers.trainer_pt_utils import get_length_grouped_indices
137139
idx_range_total = get_length_grouped_indices(

swift/megatron/trainers/gkd_trainer.py

Lines changed: 1 addition & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@
55
import torch.nn.functional as F
66
from contextlib import contextmanager
77
from functools import partial
8-
from mcore_bridge import set_random_seed
98
from megatron.core import mpu
109
from transformers.utils import ContextManagers
1110
from typing import Dict, List, Optional
@@ -167,20 +166,7 @@ def _init_resample_data_iterator(self, train_dataset):
167166
"""
168167
args = self.args
169168
resample_seed = getattr(args, 'seed', 42) + 1
170-
try:
171-
set_random_seed(
172-
resample_seed,
173-
args.data_parallel_random_init,
174-
args.te_rng_tracker,
175-
)
176-
resample_data_iterator = self._prepare_data_iterator(train_dataset, use_origin_cyclic=True)[0]
177-
finally:
178-
set_random_seed(
179-
args.seed,
180-
args.data_parallel_random_init,
181-
args.te_rng_tracker,
182-
)
183-
return resample_data_iterator
169+
return self._prepare_data_iterator(train_dataset, use_origin_cyclic=True, seed=resample_seed)[0]
184170

185171
def resample_encode_failed_inputs(self, inputs: List[Dict], max_resample_rounds: int = 10) -> List[Dict]:
186172
"""Attempt to encode each input. If encoding fails, resample until we have enough valid samples.

swift/megatron/trainers/grpo_trainer.py

Lines changed: 2 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@
55
from contextlib import contextmanager
66
from copy import copy, deepcopy
77
from functools import partial
8-
from mcore_bridge import set_random_seed
98
from megatron.core import mpu
109
from typing import Any, Dict, List, Optional, Tuple, Union
1110

@@ -209,21 +208,8 @@ def _init_resample_data_iterator(self, train_dataset):
209208
"""
210209
args = self.args
211210
resample_seed = getattr(args, 'seed', 42) + 1
212-
try:
213-
set_random_seed(
214-
resample_seed,
215-
args.data_parallel_random_init,
216-
args.te_rng_tracker,
217-
)
218-
# TODO: VPP (Virtual Pipeline Parallelism)
219-
resample_data_iterator = self._prepare_data_iterator(train_dataset, use_origin_cyclic=True)[0]
220-
finally:
221-
set_random_seed(
222-
args.seed,
223-
args.data_parallel_random_init,
224-
args.te_rng_tracker,
225-
)
226-
return resample_data_iterator
211+
# TODO: VPP (Virtual Pipeline Parallelism)
212+
return self._prepare_data_iterator(train_dataset, use_origin_cyclic=True, seed=resample_seed)[0]
227213

228214
def _build_rollout_buffer(self, data_iterator):
229215
num_gen_steps = self.steps_per_generation if self.unwrapped_models[0].training else 1

swift/megatron/trainers/utils.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -334,7 +334,7 @@ def group(self):
334334
return mpu.get_data_parallel_group()
335335

336336

337-
def build_streaming_dataloader(args, dataset, collate_fn):
337+
def build_streaming_dataloader(args, dataset, collate_fn, generator=None):
338338
base_dataloader = torch.utils.data.DataLoader(
339339
dataset,
340340
num_workers=args.dataloader_num_workers,
@@ -343,6 +343,7 @@ def build_streaming_dataloader(args, dataset, collate_fn):
343343
batch_size=args.micro_batch_size,
344344
prefetch_factor=args.dataloader_prefetch_factor if args.dataloader_num_workers > 0 else None,
345345
persistent_workers=args.dataloader_persistent_workers if args.dataloader_num_workers > 0 else False,
346+
generator=generator,
346347
)
347348
return MegatronDataLoaderDispatcher(base_dataloader)
348349

swift/megatron/utils/megatron_lm_utils.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -257,7 +257,7 @@ def save_mcore_checkpoint(
257257
if output_dir is None:
258258
output_dir = args.output_dir
259259
models = unwrap_model(models)
260-
rng_state = _get_rng_state() if models else None
260+
rng_state = None if args.no_save_rng else _get_rng_state()
261261
checkpoint_dir = os.path.join(output_dir, f'iter_{iteration:07d}')
262262
sharded_sd_metadata = get_sharded_sd_metadata(args)
263263
os.makedirs(checkpoint_dir, exist_ok=True)
@@ -284,7 +284,7 @@ def save_mcore_checkpoint(
284284
)
285285
kwargs = {'content_metadata': sharded_sd_metadata}
286286
async_save = args.async_save
287-
if not models: # save GPU memory
287+
if not models and rng_state is None: # save GPU memory when only common state remains
288288
assert 'optimizer' not in state_dict
289289
async_save = False
290290
common_path = os.path.join(checkpoint_dir, 'common.pt')

0 commit comments

Comments
 (0)