diff --git a/smauglab/configs/all_augmentations.json b/smauglab/configs/all_augmentations.json index 9c3d28f..471fb50 100644 --- a/smauglab/configs/all_augmentations.json +++ b/smauglab/configs/all_augmentations.json @@ -219,7 +219,7 @@ "keepdim": true, "p": 1.0, "p_batch": 1.0, - "retain_stats": false, + "retain_stats": true, "same_on_batch": false, "std_noise_range": [ 0.1, diff --git a/smauglab/transforms/gpu/fromSeg.py b/smauglab/transforms/gpu/fromSeg.py index 8751dc7..77bb79e 100644 --- a/smauglab/transforms/gpu/fromSeg.py +++ b/smauglab/transforms/gpu/fromSeg.py @@ -106,13 +106,31 @@ class RandomRedistributeSegGPU(ImageOnlyTransform): Mirrors the CPU `RedistributeTransform` behavior using GPU-friendly ops. Works with inputs shaped [N, C, H, W] or [N, C, D, H, W]. + + `retain_stats` defaults to True, and wants to stay that way. The whole method + operates in a per-sample [0, 1] min-max space and adds a perturbation of up to + 2.0 *in that space* -- twice the input's full dynamic range. With + `retain_stats=False` nothing maps the result back, so the output is + independent of the input's scale entirely: a z-scored patch, the same patch + scaled by 100, and anything else all come out in the same [0.5, 2.9] band with + mean 2.2. For a network fed z-scored patches that is finite, silent and badly + out of distribution -- the same defect the no-foreground branch below carries a + comment about. + + Mapping back is not a fix on its own: the perturbation is defined in + normalised units, so rescaling it by the input range makes it far larger + (measured, mean 18.4 rather than 2.2). Bounding the amplitude in input units + would be a redesign of an augmentation inherited from totalspineseg, so this + only changes which setting you get by default. `transform_params_hybrid.json` + and `transform_params_hybrid_TAGE.json` ask for False explicitly and are + unaffected. """ def __init__( self, in_seg: float = 0.2, apply_to_channel: Sequence[int] = (0,), - retain_stats: bool = False, + retain_stats: bool = True, same_on_batch: bool = False, p: float = 1.0, p_batch: float = 1.0, diff --git a/unit_tests/test_redistribute_scale.py b/unit_tests/test_redistribute_scale.py new file mode 100644 index 0000000..5f92b70 --- /dev/null +++ b/unit_tests/test_redistribute_scale.py @@ -0,0 +1,80 @@ +"""`RandomRedistributeSegGPU` must not throw away the input's intensity scale. + +The transform works in a per-sample [0, 1] min-max space and adds a perturbation +of up to 2.0 *in that space*. With `retain_stats=False` nothing maps the result +back, so the output does not depend on the input's scale at all -- a z-scored +patch and the same patch multiplied by 100 come out identical. For a network fed +z-scored patches that is finite, silent and badly out of distribution, which is +the same failure the no-foreground branch of this method carries a comment +about. + +Mapping back is not a fix on its own: the perturbation is defined in normalised +units, so rescaling it by the input range makes it much larger (measured, an +output mean of 18.4 rather than 2.2). Bounding the amplitude in input units +would be a redesign of an augmentation inherited from totalspineseg. So the +default changed and the limitation is pinned here rather than papered over. +""" + +from __future__ import annotations + +import torch + +from smauglab.transforms.gpu.fromSeg import RandomRedistributeSegGPU +from unit_tests.helpers import SmaugLabTestCase + +SHAPE = (1, 1, 16, 16, 16) + + +def seg() -> torch.Tensor: + mask = torch.zeros(*SHAPE) + mask[:, :, 4:12, 4:12, 4:12] = 1.0 + return mask + + +def run(image, **kwargs): + torch.manual_seed(3) + transform = RandomRedistributeSegGPU(p=1.0, **kwargs) + return transform.apply_transform(image.clone(), {"seg": seg()}, transform.flags) + + +class TestScaleIsPreservedByDefault(SmaugLabTestCase): + def test_the_default_keeps_the_input_statistics(self): + torch.manual_seed(0) + image = torch.randn(*SHAPE) * 1.5 + + out = run(image) + + self.assertAlmostEqual(float(out.mean()), float(image.mean()), places=4) + self.assertAlmostEqual(float(out.std()), float(image.std()), places=3) + + def test_the_default_still_changes_the_image(self): + torch.manual_seed(0) + image = torch.randn(*SHAPE) * 1.5 + + out = run(image) + + self.assertFalse(bool(torch.equal(out, image))) + self.assertTrue(bool(torch.isfinite(out).all())) + + def test_the_default_tracks_the_input_scale(self): + """Two inputs differing only in scale must not produce the same output.""" + torch.manual_seed(0) + image = torch.randn(*SHAPE) + + small = run(image) + large = run(image * 100.0) + + self.assertGreater(float(large.std() / small.std()), 50.0) + + +class TestRetainStatsFalseIsStillScaleBlind(SmaugLabTestCase): + """The known limitation, pinned so it cannot be forgotten or silently change.""" + + def test_the_output_is_independent_of_the_input_scale(self): + torch.manual_seed(0) + image = torch.randn(*SHAPE) + + outputs = [run(image * scale, retain_stats=False) for scale in (1.0, 1.5, 100.0)] + + for other in outputs[1:]: + torch.testing.assert_close(outputs[0], other)