Skip to content

Commit bfbaacd

Browse files
committed
fix(data): make ShuffleBuffer sharding explicit
Signed-off-by: kyinhub <kevinpyin@gmail.com>
1 parent 4000d9a commit bfbaacd

2 files changed

Lines changed: 79 additions & 15 deletions

File tree

monai/data/iterable_dataset.py

Lines changed: 18 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,11 @@ class ShuffleBuffer(Randomizable, IterableDataset):
7575
every iter() call, refer to the PyTorch idea:
7676
https://github.com/pytorch/pytorch/blob/v1.10.0/torch/utils/data/distributed.py#L98.
7777
epochs: number of epochs to iterate over the dataset, default to 1, -1 means infinite epochs.
78+
source_shards_by_worker: whether ``data`` already partitions its stream
79+
using ``torch.utils.data.get_worker_info``. ``None`` automatically
80+
recognizes MONAI ``IterableDataset`` sources, ``True`` avoids a
81+
second worker partition for any worker-aware source, and ``False``
82+
preserves the outer partition for unsharded iterable datasets.
7883
7984
Note:
8085
Both ``monai.data.DataLoader`` and ``torch.utils.data.DataLoader`` do not seed this class (as a subclass of
@@ -97,11 +102,22 @@ def run():
97102
98103
"""
99104

100-
def __init__(self, data, transform=None, buffer_size: int = 512, seed: int = 0, epochs: int = 1) -> None:
105+
def __init__(
106+
self,
107+
data,
108+
transform=None,
109+
buffer_size: int = 512,
110+
seed: int = 0,
111+
epochs: int = 1,
112+
source_shards_by_worker: bool | None = None,
113+
) -> None:
101114
super().__init__(data=data, transform=transform)
102115
self.size = buffer_size
103116
self.seed = seed
104117
self.epochs = epochs
118+
self.source_shards_by_worker = (
119+
isinstance(data, IterableDataset) if source_shards_by_worker is None else source_shards_by_worker
120+
)
105121
self._idx = 0
106122

107123
def randomized_pop(self, buffer):
@@ -124,17 +140,13 @@ def generate_item(self):
124140
def __iter__(self):
125141
"""Randomly pop buffered items from ``self.data``.
126142
127-
MONAI ``IterableDataset`` sources retain their existing worker
128-
partition; other sources are partitioned after shuffling.
129-
130143
Yields:
131144
Items from the shuffled source after applying the optional transform.
132145
"""
133146
self.seed += 1
134147
super().set_random_state(seed=self.seed) # make all workers in sync
135148
for _ in range(self.epochs) if self.epochs >= 0 else iter(int, 1):
136-
if isinstance(self.data, IterableDataset):
137-
# MONAI IterableDataset subclasses already partition their source per worker.
149+
if self.source_shards_by_worker:
138150
for item in self.generate_item():
139151
if self.transform is not None:
140152
item = apply_transform(self.transform, item)

tests/data/test_shuffle_buffer.py

Lines changed: 61 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -17,11 +17,33 @@
1717
from unittest.mock import patch
1818

1919
import numpy as np
20+
from torch.utils.data import IterableDataset as TorchIterableDataset
2021

2122
from monai.data import DataLoader, IterableDataset, ShuffleBuffer
23+
from monai.data import iterable_dataset as iterable_dataset_module
2224
from monai.utils import convert_data_type
2325

2426

27+
class _UnshardedMonaiIterable(IterableDataset):
28+
"""MONAI iterable subclass that intentionally does not partition itself."""
29+
30+
def __iter__(self):
31+
yield from self.data
32+
33+
34+
class _WorkerShardedTorchIterable(TorchIterableDataset):
35+
"""PyTorch iterable source that partitions itself across workers."""
36+
37+
def __init__(self, size):
38+
self.size = size
39+
40+
def __iter__(self):
41+
worker_info = iterable_dataset_module.get_worker_info()
42+
num_workers = worker_info.num_workers if worker_info is not None else 1
43+
worker_id = worker_info.id if worker_info is not None else 0
44+
yield from range(worker_id, self.size, num_workers)
45+
46+
2547
class TestShuffleBuffer(unittest.TestCase):
2648
def test_shape(self):
2749
buffer = ShuffleBuffer([1, 2, 3, 4], seed=0)
@@ -39,17 +61,47 @@ def test_shape(self):
3961
np.testing.assert_allclose(output, [[2, 3], [1, 4]], err_msg=f"seed {buffer.seed}")
4062
np.testing.assert_allclose(output2, [[1, 4], [2, 3]], err_msg=f"seed {buffer.seed}")
4163

42-
def test_iterable_dataset_is_not_sharded_twice(self):
43-
"""Verify a MONAI iterable source is partitioned exactly once."""
44-
worker_info = SimpleNamespace(num_workers=2, id=0)
45-
source = IterableDataset(range(40))
46-
buffer = ShuffleBuffer(source, transform=lambda item: item + 40, buffer_size=8, seed=7)
64+
def test_monai_iterable_source_is_detected_as_worker_sharded(self):
65+
"""Verify MONAI iterable sources avoid a second worker partition by default."""
66+
outputs = []
67+
for worker_id in range(2):
68+
source = IterableDataset(range(40))
69+
buffer = ShuffleBuffer(source, buffer_size=8, seed=7)
70+
worker_info = SimpleNamespace(num_workers=2, id=worker_id)
71+
with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info):
72+
outputs.extend(buffer)
73+
74+
self.assertEqual(len(outputs), 40)
75+
self.assertEqual(set(outputs), set(range(40)))
76+
77+
def test_worker_sharded_source_is_not_sharded_twice(self):
78+
"""Verify an explicitly worker-sharded source is not repartitioned."""
79+
sources = [IterableDataset(range(40)), _WorkerShardedTorchIterable(40)]
80+
for source in sources:
81+
outputs = []
82+
for worker_id in range(2):
83+
buffer = ShuffleBuffer(
84+
source, transform=lambda item: item + 40, buffer_size=8, seed=7, source_shards_by_worker=True
85+
)
86+
worker_info = SimpleNamespace(num_workers=2, id=worker_id)
87+
with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info):
88+
outputs.extend(buffer)
89+
90+
self.assertEqual(len(outputs), 40)
91+
self.assertEqual(set(outputs), set(range(40, 80)))
4792

48-
with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info):
49-
output = list(buffer)
93+
def test_unsharded_source_keeps_outer_worker_partition(self):
94+
"""Verify the default preserves partitioning for unsharded sources."""
95+
outputs = []
96+
for worker_id in range(2):
97+
source = _UnshardedMonaiIterable(range(40))
98+
buffer = ShuffleBuffer(source, buffer_size=8, seed=7, source_shards_by_worker=False)
99+
worker_info = SimpleNamespace(num_workers=2, id=worker_id)
100+
with patch("monai.data.iterable_dataset.get_worker_info", return_value=worker_info):
101+
outputs.extend(buffer)
50102

51-
self.assertEqual(len(output), 20)
52-
self.assertEqual(set(output), set(range(40, 80, 2)))
103+
self.assertEqual(len(outputs), 40)
104+
self.assertEqual(set(outputs), set(range(40)))
53105

54106
def test_epochs(self):
55107
buffer = ShuffleBuffer([1, 2, 3, 4], seed=0, epochs=2)

0 commit comments

Comments
 (0)