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
1 change: 1 addition & 0 deletions swift/megatron/arguments/megatron_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -514,6 +514,7 @@ class MegatronArguments(RLHFMegatronArgumentsMixin, MegatronTunerMixin):
seed: int = 42
train_dataloader_shuffle: bool = True
dataloader_num_workers: int = 4
dataloader_multiprocessing_context: Optional[Literal['fork', 'spawn', 'forkserver']] = None
dataloader_pin_memory: bool = True
dataloader_persistent_workers: bool = False
dataloader_prefetch_factor: int = 2
Expand Down
1 change: 1 addition & 0 deletions swift/megatron/trainers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -1014,6 +1014,7 @@ def _create_dataloader(self, dataset, batch_sampler):
pin_memory=args.dataloader_pin_memory,
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,
multiprocessing_context=args.dataloader_multiprocessing_context,
collate_fn=self.data_collator,
)
return dataloader
Expand Down
2 changes: 1 addition & 1 deletion swift/megatron/trainers/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -343,7 +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,
)
multiprocessing_context=args.dataloader_multiprocessing_context)
return MegatronDataLoaderDispatcher(base_dataloader)


Expand Down
1 change: 1 addition & 0 deletions swift/rlhf_trainers/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1545,6 +1545,7 @@ def get_chord_sft_dataloader(trainer,
'batch_size': batch_size,
'collate_fn': data_collator,
'num_workers': trainer.args.dataloader_num_workers,
'multiprocessing_context': trainer.args.dataloader_multiprocessing_context,
'pin_memory': trainer.args.dataloader_pin_memory,
'persistent_workers': trainer.args.dataloader_persistent_workers,
}
Expand Down
4 changes: 4 additions & 0 deletions swift/trainers/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,9 @@ class TrainArgumentsMixin:
to ['tensorboard']. If you specify `--report_to wandb`, you can set the project name through `WANDB_PROJECT`
and specify the API KEY corresponding to your account through `WANDB_API_KEY`.
dataloader_num_workers (Optional[int]): The number of subprocesses to use for data loading. Defaults to None.
dataloader_multiprocessing_context (Optional[Literal['fork', 'spawn', 'forkserver']]): The multiprocessing
context to use for data loading workers. If None, the default multiprocessing context is used. Defaults to
None.
dataloader_persistent_workers (bool): If True, the data loader workers will not be shut down after a dataset
has been consumed once. Defaults to False.
dataloader_prefetch_factor (Optional[int]): The number of batches loaded in advance by each worker. Defaults
Expand Down Expand Up @@ -148,6 +151,7 @@ class TrainArgumentsMixin:
lr_scheduler_kwargs: Optional[Union[dict, str]] = None
report_to: List[str] = field(default_factory=lambda: ['tensorboard'])
dataloader_num_workers: Optional[int] = None
dataloader_multiprocessing_context: Optional[Literal['fork', 'spawn', 'forkserver']] = None
dataloader_persistent_workers: bool = False
dataloader_prefetch_factor: Optional[int] = None
use_liger_kernel: bool = False
Expand Down
14 changes: 13 additions & 1 deletion swift/trainers/mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@
except ImportError:
sort_checkpoints = None
from types import MethodType
from typing import Callable, Dict, List, Optional
from typing import Callable, Dict, List, Optional, override

from swift.callbacks import callbacks_map
from swift.dataloader import BatchSamplerShard, DataLoaderDispatcher, DataLoaderShard
Expand Down Expand Up @@ -1266,6 +1266,7 @@ def get_sp_dataloader(self, dataset, batch_size, skip_batches=0):
'num_workers': self.args.dataloader_num_workers,
'pin_memory': self.args.dataloader_pin_memory,
'persistent_workers': self.args.dataloader_persistent_workers,
'multiprocessing_context': self.args.dataloader_multiprocessing_context
}

if not isinstance(dataset, torch.utils.data.IterableDataset):
Expand All @@ -1284,6 +1285,7 @@ def get_sp_dataloader(self, dataset, batch_size, skip_batches=0):
'num_workers': self.args.dataloader_num_workers,
'pin_memory': self.args.dataloader_pin_memory,
'persistent_workers': self.args.dataloader_persistent_workers,
'multiprocessing_context': self.args.dataloader_multiprocessing_context,
'prefetch_factor': self.args.dataloader_prefetch_factor
}
if dist.is_initialized() and dataloader_params['prefetch_factor']:
Expand All @@ -1309,6 +1311,7 @@ def get_train_dataloader(self, skip_batches=0):
'num_workers': args.dataloader_num_workers,
'pin_memory': args.dataloader_pin_memory,
'persistent_workers': args.dataloader_persistent_workers,
'multiprocessing_context': args.dataloader_multiprocessing_context,
'prefetch_factor': args.dataloader_prefetch_factor
}
batch_sampler_params = {
Expand Down Expand Up @@ -1353,6 +1356,15 @@ def _disable_group_by_length(self):
finally:
self.args.group_by_length = group_by_length

@override
def _get_dataloader(self, *args, **kwargs):
# Fix multiprocessing context here since Transformers doesn't (yet—recent main does) provide an option
dataloader = super()._get_dataloader(*args, **kwargs)
base = getattr(dataloader, 'base_dataloader', dataloader)
if (hasattr(base, 'multiprocessing_context') and base.multiprocessing_context is None):
base.multiprocessing_context = self.args.dataloader_multiprocessing_context
return dataloader

def get_eval_dataloader(self, eval_dataset=None):
dataloader = None
if self.template.sequence_parallel_size > 1:
Expand Down