From a4f38e1816236668787d610f5fcc2916d19f868c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Hendrik=20M=C3=B6ller?= Date: Thu, 24 Sep 2026 09:51:50 +0000 Subject: [PATCH] fix: draw the CPU transforms' randomness from torch, not Python's random `smauglab/transforms/rng.py` exists because `torch.manual_seed` does not seed Python's `random`: a transform reaching for `random.choice` makes a "seeded" run unreproducible, and under DDP each rank draws something different. That was fixed on the GPU side and left in place on the CPU side -- `ShapeTransform` picked its crop shape with `random.randint` and `torchio_ops.select` picked the surviving artifact with `random.choice` -- while the batchgeneratorsv2 RandomTransform wrapping both gates on `torch.rand`. One pipeline, two RNG streams. CI could not see it: `helpers.seed_everything` seeds torch, numpy *and* random, which training does not. The new tests seed only torch, as nnUNetTrainer does. `shape_min` is also clamped to the axis length; `random.randint(shape_min, s)` raised when a config asked for a minimum larger than an axis. Co-Authored-By: Claude Opus 5 --- smauglab/transforms/cpu/spatial.py | 13 ++- smauglab/transforms/cpu/torchio_ops.py | 5 +- unit_tests/test_cpu_rng.py | 105 +++++++++++++++++++++++++ 3 files changed, 118 insertions(+), 5 deletions(-) create mode 100644 unit_tests/test_cpu_rng.py diff --git a/smauglab/transforms/cpu/spatial.py b/smauglab/transforms/cpu/spatial.py index 839f7e6..da44c75 100644 --- a/smauglab/transforms/cpu/spatial.py +++ b/smauglab/transforms/cpu/spatial.py @@ -1,5 +1,3 @@ -import random - import torch import torchio as tio from batchgeneratorsv2.transforms.base.basic_transform import BasicTransform, ImageOnlyTransform @@ -81,7 +79,16 @@ def apply(self, data_dict: dict, **params) -> dict: def _apply_to_image(self, img: torch.Tensor, seg: torch.Tensor, **params) -> tuple[torch.Tensor, torch.Tensor]: # Compute random shape img_shape = img.shape[1:] - new_shape = [random.randint(params["shape_min"], s) if i not in params["ignore_axes"] else s for i, s in enumerate(img_shape)] + # torch, not random.randint: `torch.manual_seed` does not reach Python's + # `random`, so a seeded training run was not reproducible here -- which is + # exactly what smauglab.transforms.rng exists to fix, and the + # batchgeneratorsv2 RandomTransform wrapping this one already draws from + # torch. `shape_min` is clamped so a config larger than an axis crops to + # the axis instead of raising. + new_shape = [ + s if i in params["ignore_axes"] else int(torch.randint(min(params["shape_min"], s), s + 1, (1,)).item()) + for i, s in enumerate(img_shape) + ] # Find image center img_center = [s // 2 for s in img_shape] diff --git a/smauglab/transforms/cpu/torchio_ops.py b/smauglab/transforms/cpu/torchio_ops.py index 64ccdfa..c3b7c0f 100644 --- a/smauglab/transforms/cpu/torchio_ops.py +++ b/smauglab/transforms/cpu/torchio_ops.py @@ -15,13 +15,14 @@ from __future__ import annotations import gc -import random from collections.abc import Callable, Mapping from typing import cast import torch import torchio as tio +from smauglab.transforms.rng import shared_choice + #: A no-argument factory, so each call builds a freshly seeded torchio transform #: rather than reusing one instance's sampling state across the run. TransformFactory = Callable[[], tio.Transform] @@ -94,6 +95,6 @@ def select(flags: Mapping[str, bool], random_pick: bool) -> dict[str, bool]: chosen = dict(flags) enabled = [name for name, on in flags.items() if on] if random_pick and enabled: - keep = random.choice(enabled) + keep = shared_choice(enabled) chosen = {name: name == keep for name in flags} return chosen diff --git a/unit_tests/test_cpu_rng.py b/unit_tests/test_cpu_rng.py new file mode 100644 index 0000000..eec155a --- /dev/null +++ b/unit_tests/test_cpu_rng.py @@ -0,0 +1,105 @@ +"""The CPU transforms must draw from the RNG a training run actually seeds. + +`smauglab/transforms/rng.py` exists because `torch.manual_seed(...)` does not +seed Python's `random`, so a transform reaching for `random.choice` made a +"seeded" run unreproducible -- and under DDP each rank drew something different. +That was fixed on the GPU side and left in place on the CPU side: + +* `cpu/spatial.py` picked its crop shape with `random.randint` +* `cpu/torchio_ops.py` picked which artifact to keep with `random.choice` + +while the batchgeneratorsv2 `RandomTransform` wrapping both of them gates on +`torch.rand`. One pipeline, two RNG streams. + +CI could not see it: `helpers.seed_everything` seeds torch, numpy *and* random, +which training does not. These tests seed **only** torch, which is what +`nnUNetTrainer` does. +""" + +from __future__ import annotations + +import ast +import unittest +from pathlib import Path + +import torch + +from smauglab.transforms.cpu.spatial import ShapeTransform +from smauglab.transforms.cpu.torchio_ops import select + +PACKAGE = Path(__file__).resolve().parent.parent / "smauglab" + + +class TestOnlyTorchSeedingIsNeeded(unittest.TestCase): + """Seed torch alone -- as a training run does -- and expect reproducibility.""" + + def _crop_shapes(self, draws=6): + shapes = [] + torch.manual_seed(99) + transform = ShapeTransform(shape_min=4) + for _ in range(draws): + image = torch.rand(1, 12, 12, 12) + seg = torch.zeros(1, 12, 12, 12) + out = transform(image=image, segmentation=seg) + shapes.append(tuple(out["image"].shape)) + return shapes + + def test_the_crop_shape_is_reproducible_under_torch_seeding_alone(self): + first = self._crop_shapes() + second = self._crop_shapes() + + self.assertEqual(first, second) + + def test_the_crop_shape_actually_varies(self): + """Otherwise the test above would pass on a transform that does nothing.""" + self.assertGreater(len(set(self._crop_shapes(12))), 1) + + def test_select_is_reproducible_under_torch_seeding_alone(self): + flags = {"motion": True, "ghosting": True, "spike": True, "bias": True} + + def picks(n=8): + torch.manual_seed(5) + return [tuple(sorted(k for k, v in select(flags, random_pick=True).items() if v)) for _ in range(n)] + + self.assertEqual(picks(), picks()) + + def test_select_actually_varies(self): + flags = {"motion": True, "ghosting": True, "spike": True, "bias": True} + torch.manual_seed(5) + picks = {tuple(sorted(k for k, v in select(flags, random_pick=True).items() if v)) for _ in range(30)} + + self.assertGreater(len(picks), 1) + + def test_select_keeps_exactly_one(self): + flags = {"motion": True, "ghosting": True, "spike": False} + torch.manual_seed(5) + + chosen = select(flags, random_pick=True) + + self.assertEqual(sum(chosen.values()), 1) + self.assertFalse(chosen["spike"], "a disabled entry must not be picked") + + +class TestNoStdlibRandomLeftInTheTransforms(unittest.TestCase): + """A grep would catch a comment; this only counts real imports.""" + + def test_no_transform_module_imports_random(self): + for path in sorted((PACKAGE / "transforms").rglob("*.py")): + tree = ast.parse(path.read_text()) + imported = {alias.name for node in ast.walk(tree) if isinstance(node, ast.Import) for alias in node.names} | { + node.module for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) and node.module + } + + with self.subTest(module=str(path.relative_to(PACKAGE.parent))): + self.assertNotIn("random", imported, "use smauglab.transforms.rng instead of Python's random") + + +class TestCropHandlesAnOversizedMinimum(unittest.TestCase): + def test_a_shape_min_larger_than_the_axis_does_not_raise(self): + """`random.randint(shape_min, s)` raised when shape_min > s.""" + torch.manual_seed(0) + transform = ShapeTransform(shape_min=64) + + out = transform(image=torch.rand(1, 8, 8, 8), segmentation=torch.zeros(1, 8, 8, 8)) + + self.assertEqual(tuple(out["image"].shape), (1, 8, 8, 8))