Skip to content

Commit 58853d8

Browse files
Fix cyclic list handling in set_rnd (#9056)
Track container identities before recursion and preserve seed advancement across cyclic references. RED→GREEN: RecursionError and 42 != 43 → 9 focused tests passing. Full min-dependency suite: 9655 passed, 2492 skipped.
1 parent 6d28a5c commit 58853d8

2 files changed

Lines changed: 48 additions & 6 deletions

File tree

monai/data/utils.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -700,10 +700,14 @@ def set_rnd(obj, seed: int, _seen: set[int] | None = None) -> int:
700700
if _seen is None:
701701
_seen = set()
702702
if isinstance(obj, (tuple, list)): # ZipDataset.data is a list
703-
_seed = seed
703+
if id(obj) in _seen:
704+
return seed
705+
_seen.add(id(obj))
706+
has_randomizable = False
704707
for item in obj:
705-
_seed = set_rnd(item, seed=seed, _seen=_seen)
706-
return seed if _seed == seed else seed + 1 # return a different seed if there are randomizable items
708+
item_seed = set_rnd(item, seed=seed, _seen=_seen)
709+
has_randomizable = has_randomizable or item_seed != seed
710+
return seed + 1 if has_randomizable else seed
707711
if not hasattr(obj, "__dict__"):
708712
return seed # no attribute
709713
if id(obj) in _seen:

tests/data/test_dataloader.py

Lines changed: 41 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from parameterized import parameterized
2121

2222
from monai.data import CacheDataset, DataLoader, Dataset, ZipDataset
23+
from monai.data.utils import set_rnd
2324
from monai.transforms import Compose, DataStatsd, Randomizable, SimulateDelayd
2425
from monai.utils import convert_to_numpy, set_determinism
2526
from tests.test_utils import assert_allclose
@@ -28,6 +29,11 @@
2829

2930
TEST_CASE_2 = [[{"label": torch.as_tensor([[3], [2]])}, {"label": np.asarray([[1], [2]])}]]
3031

32+
_CYCLIC_DATASET_SIZE = 4
33+
_CYCLIC_BATCH_SIZE = 1
34+
_CYCLIC_NUM_WORKERS = 0
35+
_CYCLIC_TEST_SEED = 42
36+
3137

3238
class TestDataLoader(unittest.TestCase):
3339
def test_values(self):
@@ -117,18 +123,50 @@ def __init__(self):
117123
self.cfg = parent
118124

119125
def __len__(self):
120-
return 4
126+
return _CYCLIC_DATASET_SIZE
121127

122128
def __getitem__(self, index):
123129
return torch.tensor([index])
124130

125131

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+
126140
class TestLoaderRecursion(unittest.TestCase):
127141
def test_cyclic_reference_no_recursion(self):
128142
# Constructing the loader seeds the dataset (num_workers=0). A reference cycle in the
129143
# 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)
132170

133171

134172
if __name__ == "__main__":

0 commit comments

Comments
 (0)