|
20 | 20 | from parameterized import parameterized |
21 | 21 |
|
22 | 22 | from monai.data import CacheDataset, DataLoader, Dataset, ZipDataset |
| 23 | +from monai.data.utils import set_rnd |
23 | 24 | from monai.transforms import Compose, DataStatsd, Randomizable, SimulateDelayd |
24 | 25 | from monai.utils import convert_to_numpy, set_determinism |
25 | 26 | from tests.test_utils import assert_allclose |
|
28 | 29 |
|
29 | 30 | TEST_CASE_2 = [[{"label": torch.as_tensor([[3], [2]])}, {"label": np.asarray([[1], [2]])}]] |
30 | 31 |
|
| 32 | +_CYCLIC_DATASET_SIZE = 4 |
| 33 | +_CYCLIC_BATCH_SIZE = 1 |
| 34 | +_CYCLIC_NUM_WORKERS = 0 |
| 35 | +_CYCLIC_TEST_SEED = 42 |
| 36 | + |
31 | 37 |
|
32 | 38 | class TestDataLoader(unittest.TestCase): |
33 | 39 | def test_values(self): |
@@ -117,18 +123,50 @@ def __init__(self): |
117 | 123 | self.cfg = parent |
118 | 124 |
|
119 | 125 | def __len__(self): |
120 | | - return 4 |
| 126 | + return _CYCLIC_DATASET_SIZE |
121 | 127 |
|
122 | 128 | def __getitem__(self, index): |
123 | 129 | return torch.tensor([index]) |
124 | 130 |
|
125 | 131 |
|
| 132 | +class _SeedRecorder: |
| 133 | + def __init__(self): |
| 134 | + self.seed = None |
| 135 | + |
| 136 | + def set_random_state(self, seed): |
| 137 | + self.seed = seed |
| 138 | + |
| 139 | + |
126 | 140 | class TestLoaderRecursion(unittest.TestCase): |
127 | 141 | def test_cyclic_reference_no_recursion(self): |
128 | 142 | # Constructing the loader seeds the dataset (num_workers=0). A reference cycle in the |
129 | 143 | # dataset's attributes must not raise RecursionError while walking the object graph. |
130 | | - dataloader = DataLoader(_CyclicConfigDataset(), batch_size=1, num_workers=0, shuffle=False) |
131 | | - self.assertEqual(len(list(dataloader)), 4) |
| 144 | + dataloader = DataLoader( |
| 145 | + _CyclicConfigDataset(), batch_size=_CYCLIC_BATCH_SIZE, num_workers=_CYCLIC_NUM_WORKERS, shuffle=False |
| 146 | + ) |
| 147 | + self.assertEqual(len(list(dataloader)), _CYCLIC_DATASET_SIZE) |
| 148 | + |
| 149 | + def test_cyclic_list_reference_no_recursion(self): |
| 150 | + """Test seeding a dataset whose configuration list contains itself.""" |
| 151 | + dataset = _CyclicConfigDataset() |
| 152 | + dataset.cfg = [] |
| 153 | + dataset.cfg.append(dataset.cfg) |
| 154 | + dataloader = DataLoader(dataset, batch_size=_CYCLIC_BATCH_SIZE, num_workers=_CYCLIC_NUM_WORKERS, shuffle=False) |
| 155 | + self.assertEqual(len(list(dataloader)), _CYCLIC_DATASET_SIZE) |
| 156 | + |
| 157 | + def test_cyclic_list_preserves_seed_advancement(self): |
| 158 | + """Test a cyclic list does not erase seed advancement from an earlier item.""" |
| 159 | + dataset = _CyclicConfigDataset() |
| 160 | + nested_randomizable = _SeedRecorder() |
| 161 | + following_randomizable = _SeedRecorder() |
| 162 | + dataset.cfg = [nested_randomizable] |
| 163 | + dataset.cfg.append(dataset.cfg) |
| 164 | + dataset.following_randomizable = following_randomizable |
| 165 | + |
| 166 | + set_rnd(dataset, seed=_CYCLIC_TEST_SEED) |
| 167 | + |
| 168 | + self.assertEqual(nested_randomizable.seed, _CYCLIC_TEST_SEED) |
| 169 | + self.assertEqual(following_randomizable.seed, _CYCLIC_TEST_SEED + 1) |
132 | 170 |
|
133 | 171 |
|
134 | 172 | if __name__ == "__main__": |
|
0 commit comments