1717from unittest .mock import patch
1818
1919import numpy as np
20+ from torch .utils .data import IterableDataset as TorchIterableDataset
2021
2122from monai .data import DataLoader , IterableDataset , ShuffleBuffer
23+ from monai .data import iterable_dataset as iterable_dataset_module
2224from 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+
2547class 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