Skip to content
Closed
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
19 changes: 15 additions & 4 deletions monai/data/iterable_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
16 changes: 15 additions & 1 deletion tests/data/test_shuffle_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


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