Skip to content
Open
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
22 changes: 16 additions & 6 deletions swift/megatron/trainers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -577,8 +577,12 @@ def copy_path(src_path: str, tgt_path: str):
else:
raise ValueError(f'Source path is neither a file nor a directory: {src_path}')

def _prepare_data_iterator(self, train_dataset, val_dataset=None, use_origin_cyclic: bool = False):
train_dataloader, val_dataloader = self._prepare_dataloader(train_dataset, val_dataset)
def _prepare_data_iterator(self,
train_dataset,
val_dataset=None,
use_origin_cyclic: bool = False,
seed: Optional[int] = None):
train_dataloader, val_dataloader = self._prepare_dataloader(train_dataset, val_dataset, seed=seed)
train_data_iterator = iter(self.cyclic_iter(train_dataloader, use_origin_cyclic=use_origin_cyclic))
val_data_iterator = None
if val_dataset is not None:
Expand Down Expand Up @@ -973,11 +977,15 @@ def _aggregated_metrics(self, metrics, total_metrics):
total_metrics[key] = torch.tensor([0.0, 0.0], dtype=torch.float32, device=torch.cuda.current_device())
total_metrics[key] += val

def _prepare_dataloader(self, train_dataset, val_dataset=None):
def _prepare_dataloader(self, train_dataset, val_dataset=None, seed: Optional[int] = None):
args = self.args
val_dataloader = None
generator = None
if seed is not None:
generator = torch.Generator()
generator.manual_seed(seed)
if args.streaming:
train_dataloader = build_streaming_dataloader(args, train_dataset, self.data_collator)
train_dataloader = build_streaming_dataloader(args, train_dataset, self.data_collator, generator=generator)
if val_dataset is not None:
val_dataloader = build_streaming_dataloader(args, val_dataset, self.data_collator)
return train_dataloader, val_dataloader
Expand All @@ -991,8 +999,9 @@ def _prepare_dataloader(self, train_dataset, val_dataset=None):
data_sharding=args.data_sharding,
shuffle=args.train_dataloader_shuffle,
group_by_length=args.group_by_length,
seed=seed or 0,
)
train_dataloader = self._create_dataloader(train_dataset, train_batch_sampler)
train_dataloader = self._create_dataloader(train_dataset, train_batch_sampler, generator=generator)
if val_dataset is not None:
val_batch_sampler = MegatronPretrainingSampler(
total_samples=len(val_dataset),
Expand All @@ -1004,7 +1013,7 @@ def _prepare_dataloader(self, train_dataset, val_dataset=None):
val_dataloader = self._create_dataloader(val_dataset, val_batch_sampler)
return train_dataloader, val_dataloader

def _create_dataloader(self, dataset, batch_sampler):
def _create_dataloader(self, dataset, batch_sampler, generator=None):
args = self.args

dataloader = torch.utils.data.DataLoader(
Expand All @@ -1015,6 +1024,7 @@ def _create_dataloader(self, dataset, batch_sampler):
persistent_workers=args.dataloader_persistent_workers if args.dataloader_num_workers > 0 else False,
prefetch_factor=args.dataloader_prefetch_factor if args.dataloader_num_workers > 0 else None,
collate_fn=self.data_collator,
generator=generator,
)
return dataloader

Expand Down
6 changes: 4 additions & 2 deletions swift/megatron/trainers/batch_sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,7 @@ def __init__(
data_sharding,
shuffle: bool = True,
group_by_length: bool = False,
seed: int = 0,
):
# Keep a copy of input params for later use.
self.dataset = dataset
Expand All @@ -93,6 +94,7 @@ def __init__(
self.data_sharding = data_sharding
self.shuffle = shuffle
self.group_by_length = group_by_length
self.seed = seed
self.lengths = self.dataset['lengths'] if group_by_length else None
if self.lengths is not None:
self.lengths = [max(length) if isinstance(length, list) else length for length in self.lengths]
Expand Down Expand Up @@ -124,14 +126,14 @@ def __iter__(self):
start_idx = self.data_parallel_rank * bucket_size

g = torch.Generator()
g.manual_seed(self.epoch)
g.manual_seed(self.seed + self.epoch)
random_idx = torch.randperm(bucket_size, generator=g).tolist()
idx_range = [start_idx + x for x in random_idx[bucket_offset:]]
else:
full_bucket_size = (self.total_samples // self.micro_batch_size) * self.micro_batch_size
full_bucket_offset = current_epoch_samples
g = torch.Generator()
g.manual_seed(self.epoch)
g.manual_seed(self.seed + self.epoch)
if self.group_by_length:
from transformers.trainer_pt_utils import get_length_grouped_indices
idx_range_total = get_length_grouped_indices(
Expand Down
16 changes: 1 addition & 15 deletions swift/megatron/trainers/gkd_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
import torch.nn.functional as F
from contextlib import contextmanager
from functools import partial
from mcore_bridge import set_random_seed
from megatron.core import mpu
from transformers.utils import ContextManagers
from typing import Dict, List, Optional
Expand Down Expand Up @@ -167,20 +166,7 @@ def _init_resample_data_iterator(self, train_dataset):
"""
args = self.args
resample_seed = getattr(args, 'seed', 42) + 1
try:
set_random_seed(
resample_seed,
args.data_parallel_random_init,
args.te_rng_tracker,
)
resample_data_iterator = self._prepare_data_iterator(train_dataset, use_origin_cyclic=True)[0]
finally:
set_random_seed(
args.seed,
args.data_parallel_random_init,
args.te_rng_tracker,
)
return resample_data_iterator
return self._prepare_data_iterator(train_dataset, use_origin_cyclic=True, seed=resample_seed)[0]

def resample_encode_failed_inputs(self, inputs: List[Dict], max_resample_rounds: int = 10) -> List[Dict]:
"""Attempt to encode each input. If encoding fails, resample until we have enough valid samples.
Expand Down
18 changes: 2 additions & 16 deletions swift/megatron/trainers/grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
from contextlib import contextmanager
from copy import copy, deepcopy
from functools import partial
from mcore_bridge import set_random_seed
from megatron.core import mpu
from typing import Any, Dict, List, Optional, Tuple, Union

Expand Down Expand Up @@ -209,21 +208,8 @@ def _init_resample_data_iterator(self, train_dataset):
"""
args = self.args
resample_seed = getattr(args, 'seed', 42) + 1
try:
set_random_seed(
resample_seed,
args.data_parallel_random_init,
args.te_rng_tracker,
)
# TODO: VPP (Virtual Pipeline Parallelism)
resample_data_iterator = self._prepare_data_iterator(train_dataset, use_origin_cyclic=True)[0]
finally:
set_random_seed(
args.seed,
args.data_parallel_random_init,
args.te_rng_tracker,
)
return resample_data_iterator
# TODO: VPP (Virtual Pipeline Parallelism)
return self._prepare_data_iterator(train_dataset, use_origin_cyclic=True, seed=resample_seed)[0]

def _build_rollout_buffer(self, data_iterator):
num_gen_steps = self.steps_per_generation if self.unwrapped_models[0].training else 1
Expand Down
3 changes: 2 additions & 1 deletion swift/megatron/trainers/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -334,7 +334,7 @@ def group(self):
return mpu.get_data_parallel_group()


def build_streaming_dataloader(args, dataset, collate_fn):
def build_streaming_dataloader(args, dataset, collate_fn, generator=None):
base_dataloader = torch.utils.data.DataLoader(
dataset,
num_workers=args.dataloader_num_workers,
Expand All @@ -343,6 +343,7 @@ def build_streaming_dataloader(args, dataset, collate_fn):
batch_size=args.micro_batch_size,
prefetch_factor=args.dataloader_prefetch_factor if args.dataloader_num_workers > 0 else None,
persistent_workers=args.dataloader_persistent_workers if args.dataloader_num_workers > 0 else False,
generator=generator,
)
return MegatronDataLoaderDispatcher(base_dataloader)

Expand Down
4 changes: 2 additions & 2 deletions swift/megatron/utils/megatron_lm_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,7 @@ def save_mcore_checkpoint(
if output_dir is None:
output_dir = args.output_dir
models = unwrap_model(models)
rng_state = _get_rng_state() if models else None
rng_state = None if args.no_save_rng else _get_rng_state()
checkpoint_dir = os.path.join(output_dir, f'iter_{iteration:07d}')
sharded_sd_metadata = get_sharded_sd_metadata(args)
os.makedirs(checkpoint_dir, exist_ok=True)
Expand All @@ -277,7 +277,7 @@ def save_mcore_checkpoint(
)
kwargs = {'content_metadata': sharded_sd_metadata}
async_save = args.async_save
if not models: # save GPU memory
if not models and rng_state is None: # save GPU memory when only common state remains
assert 'optimizer' not in state_dict
async_save = False
common_path = os.path.join(checkpoint_dir, 'common.pt')
Expand Down
Loading