Skip to content
Merged
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
2 changes: 1 addition & 1 deletion docs/source/Instruction/Command-line-parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -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`。
Expand Down
2 changes: 1 addition & 1 deletion docs/source_en/Instruction/Command-line-parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down
5 changes: 4 additions & 1 deletion swift/trainers/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down
84 changes: 84 additions & 0 deletions tests/general/test_dataloader_persistent_workers.py
Original file line number Diff line number Diff line change
@@ -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()
Loading