From 4000d9ad83d47a922c10a2f655e439d13020c087 Mon Sep 17 00:00:00 2001 From: kyinhub Date: Sat, 25 Jul 2026 11:11:20 -0700 Subject: [PATCH 1/3] fix(data): avoid double sharding shuffle buffer Signed-off-by: kyinhub --- monai/data/iterable_dataset.py | 19 +++++++++++++++---- tests/data/test_shuffle_buffer.py | 16 +++++++++++++++- 2 files changed, 30 insertions(+), 5 deletions(-) diff --git a/monai/data/iterable_dataset.py b/monai/data/iterable_dataset.py index 3ee70bf69eb..0e85b214c1b 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 274121ec772..6ecc0efc43e 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)] From e525693264bb1f895fbd4333a9a6648987dc488e Mon Sep 17 00:00:00 2001 From: kyinhub Date: Sat, 25 Jul 2026 14:04:37 -0700 Subject: [PATCH 2/3] fix(data): make ShuffleBuffer sharding explicit Signed-off-by: kyinhub --- monai/data/iterable_dataset.py | 24 ++++++++--- tests/data/test_shuffle_buffer.py | 70 +++++++++++++++++++++++++++---- 2 files changed, 79 insertions(+), 15 deletions(-) diff --git a/monai/data/iterable_dataset.py b/monai/data/iterable_dataset.py index 0e85b214c1b..d88cd283c69 100644 --- a/monai/data/iterable_dataset.py +++ b/monai/data/iterable_dataset.py @@ -75,6 +75,11 @@ class ShuffleBuffer(Randomizable, IterableDataset): every iter() call, refer to the PyTorch idea: https://github.com/pytorch/pytorch/blob/v1.10.0/torch/utils/data/distributed.py#L98. epochs: number of epochs to iterate over the dataset, default to 1, -1 means infinite epochs. + source_shards_by_worker: whether ``data`` already partitions its stream + using ``torch.utils.data.get_worker_info``. ``None`` automatically + recognizes MONAI ``IterableDataset`` sources, ``True`` avoids a + second worker partition for any worker-aware source, and ``False`` + preserves the outer partition for unsharded iterable datasets. Note: Both ``monai.data.DataLoader`` and ``torch.utils.data.DataLoader`` do not seed this class (as a subclass of @@ -97,11 +102,22 @@ def run(): """ - def __init__(self, data, transform=None, buffer_size: int = 512, seed: int = 0, epochs: int = 1) -> None: + def __init__( + self, + data, + transform=None, + buffer_size: int = 512, + seed: int = 0, + epochs: int = 1, + source_shards_by_worker: bool | None = None, + ) -> None: super().__init__(data=data, transform=transform) self.size = buffer_size self.seed = seed self.epochs = epochs + self.source_shards_by_worker = ( + isinstance(data, IterableDataset) if source_shards_by_worker is None else source_shards_by_worker + ) self._idx = 0 def randomized_pop(self, buffer): @@ -124,17 +140,13 @@ def generate_item(self): def __iter__(self): """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): - if isinstance(self.data, IterableDataset): - # MONAI IterableDataset subclasses already partition their source per worker. + if self.source_shards_by_worker: for item in self.generate_item(): if self.transform is not None: item = apply_transform(self.transform, item) diff --git a/tests/data/test_shuffle_buffer.py b/tests/data/test_shuffle_buffer.py index 6ecc0efc43e..654a4928d05 100644 --- a/tests/data/test_shuffle_buffer.py +++ b/tests/data/test_shuffle_buffer.py @@ -17,11 +17,33 @@ from unittest.mock import patch import numpy as np +from torch.utils.data import IterableDataset as TorchIterableDataset from monai.data import DataLoader, IterableDataset, ShuffleBuffer +from monai.data import iterable_dataset as iterable_dataset_module from monai.utils import convert_data_type +class _UnshardedMonaiIterable(IterableDataset): + """MONAI iterable subclass that intentionally does not partition itself.""" + + def __iter__(self): + yield from self.data + + +class _WorkerShardedTorchIterable(TorchIterableDataset): + """PyTorch iterable source that partitions itself across workers.""" + + def __init__(self, size): + self.size = size + + def __iter__(self): + worker_info = iterable_dataset_module.get_worker_info() + num_workers = worker_info.num_workers if worker_info is not None else 1 + worker_id = worker_info.id if worker_info is not None else 0 + yield from range(worker_id, self.size, num_workers) + + class TestShuffleBuffer(unittest.TestCase): def test_shape(self): buffer = ShuffleBuffer([1, 2, 3, 4], seed=0) @@ -39,17 +61,47 @@ 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) + def test_monai_iterable_source_is_detected_as_worker_sharded(self): + """Verify MONAI iterable sources avoid a second worker partition by default.""" + outputs = [] + for worker_id in range(2): + source = IterableDataset(range(40)) + buffer = ShuffleBuffer(source, buffer_size=8, seed=7) + worker_info = SimpleNamespace(num_workers=2, id=worker_id) + with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info): + outputs.extend(buffer) + + self.assertEqual(len(outputs), 40) + self.assertEqual(set(outputs), set(range(40))) + + def test_worker_sharded_source_is_not_sharded_twice(self): + """Verify an explicitly worker-sharded source is not repartitioned.""" + sources = [IterableDataset(range(40)), _WorkerShardedTorchIterable(40)] + for source in sources: + outputs = [] + for worker_id in range(2): + buffer = ShuffleBuffer( + source, transform=lambda item: item + 40, buffer_size=8, seed=7, source_shards_by_worker=True + ) + worker_info = SimpleNamespace(num_workers=2, id=worker_id) + with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info): + outputs.extend(buffer) + + self.assertEqual(len(outputs), 40) + self.assertEqual(set(outputs), set(range(40, 80))) - with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info): - output = list(buffer) + def test_explicit_unsharded_source_keeps_outer_worker_partition(self): + """Verify explicit unsharded mode preserves outer worker partitioning.""" + outputs = [] + for worker_id in range(2): + source = _UnshardedMonaiIterable(range(40)) + buffer = ShuffleBuffer(source, buffer_size=8, seed=7, source_shards_by_worker=False) + worker_info = SimpleNamespace(num_workers=2, id=worker_id) + with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info): + outputs.extend(buffer) - self.assertEqual(len(output), 20) - self.assertEqual(set(output), set(range(40, 80, 2))) + self.assertEqual(len(outputs), 40) + self.assertEqual(set(outputs), set(range(40))) def test_epochs(self): buffer = ShuffleBuffer([1, 2, 3, 4], seed=0, epochs=2) From cb3a713add8ec0ec75d8dc3bd296bc748849e266 Mon Sep 17 00:00:00 2001 From: kyinhub Date: Sun, 26 Jul 2026 02:11:31 -0700 Subject: [PATCH 3/3] docs(data): document ShuffleBuffer constructor Signed-off-by: kyinhub --- monai/data/iterable_dataset.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/monai/data/iterable_dataset.py b/monai/data/iterable_dataset.py index d88cd283c69..93be96ff33b 100644 --- a/monai/data/iterable_dataset.py +++ b/monai/data/iterable_dataset.py @@ -111,6 +111,18 @@ def __init__( epochs: int = 1, source_shards_by_worker: bool | None = None, ) -> None: + """Initialize the shuffle buffer. + + Args: + data: input data source to load, shuffle, and optionally transform. + transform: a callable data transform applied to each yielded item. + buffer_size: maximum number of items stored before random popping. + seed: random seed used to initialize the worker random states. + epochs: number of source iterations, where ``-1`` means infinite. + source_shards_by_worker: whether ``data`` already partitions its + stream using ``torch.utils.data.get_worker_info``. ``None`` + automatically recognizes MONAI ``IterableDataset`` sources. + """ super().__init__(data=data, transform=transform) self.size = buffer_size self.seed = seed