diff --git a/monai/data/iterable_dataset.py b/monai/data/iterable_dataset.py index 3ee70bf69e..0e85b214c1 100644 --- a/monai/data/iterable_dataset.py +++ b/monai/data/iterable_dataset.py @@ -122,14 +122,25 @@ def generate_item(self): yield self.randomized_pop(buffer) def __iter__(self): - """ - Randomly pop buffered items from `self.data`. - Multiple dataloader workers sharing this dataset will generate identical item sequences. + """Randomly pop buffered items from ``self.data``. + + MONAI ``IterableDataset`` sources retain their existing worker + partition; other sources are partitioned after shuffling. + + Yields: + Items from the shuffled source after applying the optional transform. """ self.seed += 1 super().set_random_state(seed=self.seed) # make all workers in sync for _ in range(self.epochs) if self.epochs >= 0 else iter(int, 1): - yield from IterableDataset(self.generate_item(), transform=self.transform) + if isinstance(self.data, IterableDataset): + # MONAI IterableDataset subclasses already partition their source per worker. + for item in self.generate_item(): + if self.transform is not None: + item = apply_transform(self.transform, item) + yield item + else: + yield from IterableDataset(self.generate_item(), transform=self.transform) def randomize(self, size: int) -> None: self._idx = self.R.randint(size) diff --git a/tests/data/test_shuffle_buffer.py b/tests/data/test_shuffle_buffer.py index 274121ec77..6ecc0efc43 100644 --- a/tests/data/test_shuffle_buffer.py +++ b/tests/data/test_shuffle_buffer.py @@ -13,10 +13,12 @@ import sys import unittest +from types import SimpleNamespace +from unittest.mock import patch import numpy as np -from monai.data import DataLoader, ShuffleBuffer +from monai.data import DataLoader, IterableDataset, ShuffleBuffer from monai.utils import convert_data_type @@ -37,6 +39,18 @@ def test_shape(self): np.testing.assert_allclose(output, [[2, 3], [1, 4]], err_msg=f"seed {buffer.seed}") np.testing.assert_allclose(output2, [[1, 4], [2, 3]], err_msg=f"seed {buffer.seed}") + def test_iterable_dataset_is_not_sharded_twice(self): + """Verify a MONAI iterable source is partitioned exactly once.""" + worker_info = SimpleNamespace(num_workers=2, id=0) + source = IterableDataset(range(40)) + buffer = ShuffleBuffer(source, transform=lambda item: item + 40, buffer_size=8, seed=7) + + with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info): + output = list(buffer) + + self.assertEqual(len(output), 20) + self.assertEqual(set(output), set(range(40, 80, 2))) + def test_epochs(self): buffer = ShuffleBuffer([1, 2, 3, 4], seed=0, epochs=2) output = [convert_data_type(x, np.ndarray)[0] for x in DataLoader(dataset=buffer, batch_size=2)]