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
2 changes: 1 addition & 1 deletion smauglab/configs/all_augmentations.json
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
20 changes: 19 additions & 1 deletion smauglab/transforms/gpu/fromSeg.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
80 changes: 80 additions & 0 deletions unit_tests/test_redistribute_scale.py
Original file line number Diff line number Diff line change
@@ -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)
Loading