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
22 changes: 17 additions & 5 deletions monai/data/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -686,28 +686,40 @@ def worker_init_fn(worker_id: int) -> None:
set_rnd(worker_info.dataset, seed=worker_info.seed) # type: ignore[union-attr]


def set_rnd(obj, seed: int) -> int:
def set_rnd(obj, seed: int, _seen: set[int] | None = None) -> int:
"""
Set seed or random state for all randomizable properties of obj.

Args:
obj: object to set seed or random state for.
seed: set the random state with an integer seed.
_seen: internal set of already-visited object ids, used to guard against
infinite recursion on cyclic object graphs (e.g. OmegaConf/Hydra
configs whose child nodes back-reference their parent, see issue #8087).
"""
if _seen is None:
_seen = set()
if isinstance(obj, (tuple, list)): # ZipDataset.data is a list
_seed = seed
if id(obj) in _seen:
return seed
_seen.add(id(obj))
has_randomizable = False
for item in obj:
_seed = set_rnd(item, seed=seed)
return seed if _seed == seed else seed + 1 # return a different seed if there are randomizable items
item_seed = set_rnd(item, seed=seed, _seen=_seen)
has_randomizable = has_randomizable or item_seed != seed
return seed + 1 if has_randomizable else seed
if not hasattr(obj, "__dict__"):
return seed # no attribute
if id(obj) in _seen:
return seed # already visited: avoid infinite recursion on cyclic references
_seen.add(id(obj))
Comment thread
ousamabenyounes marked this conversation as resolved.
if hasattr(obj, "set_random_state"):
obj.set_random_state(seed=seed % MAX_SEED)
return seed + 1 # a different seed for the next component
for key in obj.__dict__:
if key.startswith("__"): # skip the private methods
continue
seed = set_rnd(obj.__dict__[key], seed=seed)
seed = set_rnd(obj.__dict__[key], seed=seed, _seen=_seen)
return seed


Expand Down
70 changes: 70 additions & 0 deletions tests/data/test_dataloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,15 @@
from __future__ import annotations

import sys
import types
import unittest

import numpy as np
import torch
from parameterized import parameterized

from monai.data import CacheDataset, DataLoader, Dataset, ZipDataset
from monai.data.utils import set_rnd
from monai.transforms import Compose, DataStatsd, Randomizable, SimulateDelayd
from monai.utils import convert_to_numpy, set_determinism
from tests.test_utils import assert_allclose
Expand All @@ -27,6 +29,11 @@

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

_CYCLIC_DATASET_SIZE = 4
_CYCLIC_BATCH_SIZE = 1
_CYCLIC_NUM_WORKERS = 0
_CYCLIC_TEST_SEED = 42


class TestDataLoader(unittest.TestCase):
def test_values(self):
Expand Down Expand Up @@ -99,5 +106,68 @@ def test_zipdataset(self):
assert_allclose(np.stack(output).flatten()[:7], np.array([594, 170, 594, 170, 594, 170, 524]))


class _CyclicConfigDataset(torch.utils.data.Dataset):
"""
Dataset holding an attribute whose object graph contains a reference cycle.

This mirrors OmegaConf/Hydra configs, whose child nodes hold a back-reference
to their parent node. Seeding such a dataset used to recurse forever in
``monai.data.utils.set_rnd`` (see issue #8087).
"""

def __init__(self):
parent = types.SimpleNamespace()
child = types.SimpleNamespace()
parent.child = child
child.parent = parent # reference cycle, as in an OmegaConf parent/child graph
self.cfg = parent

def __len__(self):
return _CYCLIC_DATASET_SIZE

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


class _SeedRecorder:
def __init__(self):
self.seed = None

def set_random_state(self, seed):
self.seed = seed


class TestLoaderRecursion(unittest.TestCase):
def test_cyclic_reference_no_recursion(self):
# Constructing the loader seeds the dataset (num_workers=0). A reference cycle in the
# dataset's attributes must not raise RecursionError while walking the object graph.
dataloader = DataLoader(
_CyclicConfigDataset(), batch_size=_CYCLIC_BATCH_SIZE, num_workers=_CYCLIC_NUM_WORKERS, shuffle=False
)
self.assertEqual(len(list(dataloader)), _CYCLIC_DATASET_SIZE)

def test_cyclic_list_reference_no_recursion(self):
"""Test seeding a dataset whose configuration list contains itself."""
dataset = _CyclicConfigDataset()
dataset.cfg = []
dataset.cfg.append(dataset.cfg)
dataloader = DataLoader(dataset, batch_size=_CYCLIC_BATCH_SIZE, num_workers=_CYCLIC_NUM_WORKERS, shuffle=False)
self.assertEqual(len(list(dataloader)), _CYCLIC_DATASET_SIZE)

def test_cyclic_list_preserves_seed_advancement(self):
"""Test a cyclic list does not erase seed advancement from an earlier item."""
dataset = _CyclicConfigDataset()
nested_randomizable = _SeedRecorder()
following_randomizable = _SeedRecorder()
dataset.cfg = [nested_randomizable]
dataset.cfg.append(dataset.cfg)
dataset.following_randomizable = following_randomizable

set_rnd(dataset, seed=_CYCLIC_TEST_SEED)

self.assertEqual(nested_randomizable.seed, _CYCLIC_TEST_SEED)
self.assertEqual(following_randomizable.seed, _CYCLIC_TEST_SEED + 1)


if __name__ == "__main__":
unittest.main()
Loading