Skip to content
Open
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
45 changes: 40 additions & 5 deletions monai/data/iterable_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -97,11 +102,34 @@ 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:
"""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
self.epochs = epochs
self.source_shards_by_worker = (
isinstance(data, IterableDataset) if source_shards_by_worker is None else source_shards_by_worker
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
self._idx = 0

def randomized_pop(self, buffer):
Expand All @@ -122,14 +150,21 @@ 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``.

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 self.source_shards_by_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)
Expand Down
68 changes: 67 additions & 1 deletion tests/data/test_shuffle_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,13 +13,37 @@

import sys
import unittest
from types import SimpleNamespace
from unittest.mock import patch

import numpy as np
from torch.utils.data import IterableDataset as TorchIterableDataset

from monai.data import DataLoader, ShuffleBuffer
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)
Expand All @@ -37,6 +61,48 @@ 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_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)))

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(outputs), 40)
self.assertEqual(set(outputs), set(range(40)))

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)]
Expand Down
Loading