diff --git a/swift/megatron/arguments/megatron_args.py b/swift/megatron/arguments/megatron_args.py index 5afaacd220..5bde2db0c9 100644 --- a/swift/megatron/arguments/megatron_args.py +++ b/swift/megatron/arguments/megatron_args.py @@ -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 diff --git a/swift/megatron/trainers/base.py b/swift/megatron/trainers/base.py index 048bf4dcf5..60ff19eec5 100644 --- a/swift/megatron/trainers/base.py +++ b/swift/megatron/trainers/base.py @@ -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 diff --git a/swift/megatron/trainers/utils.py b/swift/megatron/trainers/utils.py index 66d1e5dfd9..4a55da122a 100644 --- a/swift/megatron/trainers/utils.py +++ b/swift/megatron/trainers/utils.py @@ -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) diff --git a/swift/rlhf_trainers/utils.py b/swift/rlhf_trainers/utils.py index 88289c53c7..58ce911239 100644 --- a/swift/rlhf_trainers/utils.py +++ b/swift/rlhf_trainers/utils.py @@ -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, } diff --git a/swift/trainers/arguments.py b/swift/trainers/arguments.py index 05c3a18bb2..5147cd9d09 100644 --- a/swift/trainers/arguments.py +++ b/swift/trainers/arguments.py @@ -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 @@ -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 diff --git a/swift/trainers/mixin.py b/swift/trainers/mixin.py index 578e4a41b8..02f672d5e0 100644 --- a/swift/trainers/mixin.py +++ b/swift/trainers/mixin.py @@ -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 @@ -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): @@ -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']: @@ -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 = { @@ -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: