From 6993bd1cbe8460072f1cf2de2b9cd13533de564f Mon Sep 17 00:00:00 2001 From: "zhao.wang" <57819425+Excelius-Wang@users.noreply.github.com> Date: Mon, 14 Sep 2026 16:15:58 +0800 Subject: [PATCH] fix(train): disable persistent workers when no workers are used Signed-off-by: zhao.wang <57819425+Excelius-Wang@users.noreply.github.com> --- .../Instruction/Command-line-parameters.md | 2 +- .../Instruction/Command-line-parameters.md | 2 +- swift/trainers/arguments.py | 5 +- .../test_dataloader_persistent_workers.py | 84 +++++++++++++++++++ 4 files changed, 90 insertions(+), 3 deletions(-) create mode 100644 tests/general/test_dataloader_persistent_workers.py diff --git a/docs/source/Instruction/Command-line-parameters.md b/docs/source/Instruction/Command-line-parameters.md index 21bb644e60..4b3d43316d 100644 --- a/docs/source/Instruction/Command-line-parameters.md +++ b/docs/source/Instruction/Command-line-parameters.md @@ -295,7 +295,7 @@ ENV: - 🔥ddp_find_unused_parameters: 默认为None。 - 🔥dataloader_num_workers: 默认为None,若是windows平台,则设置为0,否则设置为1。 - dataloader_pin_memory: 默认为True。 -- dataloader_persistent_workers: 默认为False。 +- dataloader_persistent_workers: 默认为True;当 `dataloader_num_workers=0` 时自动设为False。 - dataloader_prefetch_factor: 默认为None。若 `dataloader_num_workers > 0`,则设置为2。每个工作进程预先加载的批次数量。2 表示所有工作进程总共会预取 2 * num_workers 个批次。 - train_dataloader_shuffle: CPT/SFT训练的dataloader是否随机,默认为True。该参数对IterableDataset无效(即对流式数据集失效)。IterableDataset采用顺序的方式读取。 - optim: 优化器,默认值为 `"adamw_torch"` (对于 torch>=2.8 为 `"adamw_torch_fused"`)。完整的优化器列表请参见 [training_args.py](https://github.com/huggingface/transformers/blob/main/src/transformers/training_args.py) 中的 `OptimizerNames`。 diff --git a/docs/source_en/Instruction/Command-line-parameters.md b/docs/source_en/Instruction/Command-line-parameters.md index e21b0c687a..01a019d32c 100644 --- a/docs/source_en/Instruction/Command-line-parameters.md +++ b/docs/source_en/Instruction/Command-line-parameters.md @@ -300,7 +300,7 @@ Other important parameters: - 🔥ddp_find_unused_parameters: Default is `None`. - 🔥dataloader_num_workers: Default is `None`. On Windows, set to 0; otherwise, 1. - dataloader_pin_memory: Default is `True`. -- dataloader_persistent_workers: Default is `False`. +- dataloader_persistent_workers: Default is `True`; automatically set to `False` when `dataloader_num_workers=0`. - dataloader_prefetch_factor: Default is `None`. If `dataloader_num_workers > 0`, it is set to 2. Number of batches loaded in advance by each worker. 2 means there will be a total of 2 * num_workers batches prefetched across all workers. - train_dataloader_shuffle: Whether to shuffle the dataloader for CPT/SFT training, default is True. This parameter is ineffective for IterableDataset (i.e., it doesn't work for streaming datasets). IterableDataset reads data sequentially. - optim: The optimizer, defaults to `"adamw_torch"` (for torch>=2.8 `"adamw_torch_fused"`). For a complete list of optimizers, please see `OptimizerNames` in [training_args.py](https://github.com/huggingface/transformers/blob/main/src/transformers/training_args.py). diff --git a/swift/trainers/arguments.py b/swift/trainers/arguments.py index 5036e7cd27..689ecc8a66 100644 --- a/swift/trainers/arguments.py +++ b/swift/trainers/arguments.py @@ -49,7 +49,7 @@ class TrainArgumentsMixin: 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_persistent_workers (bool): If True, the data loader workers will not be shut down after a dataset - has been consumed once. Defaults to True. + has been consumed once. Defaults to True. Disabled when dataloader_num_workers is 0. dataloader_prefetch_factor (Optional[int]): The number of batches loaded in advance by each worker. Defaults to None. use_liger_kernel (bool): Whether to use the Liger kernel for optimization. Defaults to False. @@ -304,6 +304,9 @@ def __post_init__(self): else: self.dataloader_num_workers = 1 logger.info(f'Setting args.dataloader_num_workers: {self.dataloader_num_workers}') + if self.dataloader_num_workers == 0 and self.dataloader_persistent_workers: + self.dataloader_persistent_workers = False + logger.info('Setting args.dataloader_persistent_workers: False because dataloader_num_workers is 0.') if self.dataloader_prefetch_factor is None and self.dataloader_num_workers > 0: self.dataloader_prefetch_factor = 2 if self.eval_use_evalscope: diff --git a/tests/general/test_dataloader_persistent_workers.py b/tests/general/test_dataloader_persistent_workers.py new file mode 100644 index 0000000000..014ea37276 --- /dev/null +++ b/tests/general/test_dataloader_persistent_workers.py @@ -0,0 +1,84 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import tempfile +import torch +import unittest +from datasets import Dataset +from transformers import Trainer as HfTrainer +from transformers import default_data_collator +from types import SimpleNamespace +from unittest import mock + +from swift.trainers import Seq2SeqTrainingArguments, TrainingArguments +from swift.trainers.mixin import DataLoaderMixin + + +class TestDataloaderPersistentWorkers(unittest.TestCase): + + def setUp(self): + self.tmp_dir = tempfile.TemporaryDirectory() + self.addCleanup(self.tmp_dir.cleanup) + + def make_args(self, args_cls, **kwargs): + return args_cls( + output_dir=self.tmp_dir.name, use_cpu=True, report_to=[], train_dataloader_shuffle=False, **kwargs) + + def check_training_loader(self, args, dataset): + trainer = DataLoaderMixin() + trainer.args = args + trainer.template = SimpleNamespace(sequence_parallel_size=1) + trainer.accelerator = SimpleNamespace(device=torch.device('cpu')) + trainer.train_dataset = dataset + trainer.data_collator = default_data_collator + trainer._train_batch_size = 1 + loader = trainer.get_train_dataloader() + self.assertEqual([batch['input_ids'].item() for batch in loader], [1, 2]) + + def test_zero_workers_can_load_map_and_iterable_datasets(self): + rows = Dataset.from_dict({'input_ids': [[1], [2]]}) + for args_cls in (TrainingArguments, Seq2SeqTrainingArguments): + for dataset in (rows, rows.to_iterable_dataset()): + with self.subTest(args_cls=args_cls.__name__, dataset=type(dataset).__name__): + args = self.make_args(args_cls, dataloader_num_workers=0) + self.check_training_loader(args, dataset) + + def test_zero_workers_can_load_eval_and_prediction_data(self): + rows = Dataset.from_dict({'input_ids': [[1], [2]]}) + for args_cls in (TrainingArguments, Seq2SeqTrainingArguments): + with self.subTest(args_cls=args_cls.__name__): + args = self.make_args(args_cls, dataloader_num_workers=0, remove_unused_columns=False) + trainer = HfTrainer( + model=torch.nn.Linear(1, 1), args=args, eval_dataset=rows, data_collator=default_data_collator) + for loader in (trainer.get_eval_dataloader(), trainer.get_test_dataloader(rows)): + self.assertEqual(torch.cat([batch['input_ids'] for batch in loader]).flatten().tolist(), [1, 2]) + + def test_explicit_persistence_settings_with_zero_workers(self): + for persistent in (True, False): + with self.subTest(persistent=persistent): + args = self.make_args( + Seq2SeqTrainingArguments, dataloader_num_workers=0, dataloader_persistent_workers=persistent) + self.check_training_loader(args, [{'input_ids': [1]}, {'input_ids': [2]}]) + + def test_explicit_prefetch_with_zero_workers_still_raises(self): + with self.assertRaises(ValueError): + self.make_args(Seq2SeqTrainingArguments, dataloader_num_workers=0, dataloader_prefetch_factor=2) + + def test_windows_default_can_load_data(self): + # Exercise Windows argument defaults without starting platform-specific worker processes. + for args_cls in (TrainingArguments, Seq2SeqTrainingArguments): + with self.subTest(args_cls=args_cls.__name__): + with mock.patch('swift.trainers.arguments.platform.system', return_value='Windows'): + args = self.make_args(args_cls) + self.assertEqual(args.dataloader_num_workers, 0) + self.check_training_loader(args, [{'input_ids': [1]}, {'input_ids': [2]}]) + + def test_positive_workers_keep_default_and_explicit_opt_out(self): + for args_cls in (TrainingArguments, Seq2SeqTrainingArguments): + with self.subTest(args_cls=args_cls.__name__): + args = self.make_args(args_cls, dataloader_num_workers=1) + self.assertTrue(args.dataloader_persistent_workers) + args = self.make_args(args_cls, dataloader_num_workers=1, dataloader_persistent_workers=False) + self.assertFalse(args.dataloader_persistent_workers) + + +if __name__ == '__main__': + unittest.main()