From b39b9396d7996727f614e8733a69fb7611fbda0c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Hendrik=20M=C3=B6ller?= Date: Thu, 24 Sep 2026 10:00:22 +0000 Subject: [PATCH] fix: deform the SynthSeg label map before cropping it The module docstring states the order as "spatial deform -> [optional random crop] -> flip", and lab2im's RandomSpatialDeformation(output_shape=...) resamples the full label map onto the output grid. The code cropped first. That is not a stale comment. `warp_volume` pads with zeros, so deforming an already-cropped box pulls background in from immediately outside it, where the reference pulls in the real anatomy that sits there. Measured over 40 seeds on a fully labelled 24-cube cropped to 16 -- every background voxel in the output is therefore invented padding: crop -> deform: mean 13.5% background, 0/40 draws with none deform -> crop: mean 3.3% background, 7/40 draws with none Deforming first does not remove the effect entirely (a crop landing against the volume edge still meets padding), but it is the order the reference uses and it cuts the imported background fourfold. No shipped config changes: `output_shape` defaults to None and is not among RandomSynthSegGPU's accepted parameters, so the crop only runs for a caller constructing SynthSegGenerator directly. Both facts are asserted in the test rather than left as a claim. Co-Authored-By: Claude Opus 5 --- smauglab/transforms/synthseg/generator.py | 20 +++- unit_tests/test_synthseg_crop_order.py | 117 ++++++++++++++++++++++ 2 files changed, 132 insertions(+), 5 deletions(-) create mode 100644 unit_tests/test_synthseg_crop_order.py diff --git a/smauglab/transforms/synthseg/generator.py b/smauglab/transforms/synthseg/generator.py index 22fcf81..9626eff 100644 --- a/smauglab/transforms/synthseg/generator.py +++ b/smauglab/transforms/synthseg/generator.py @@ -242,11 +242,7 @@ def forward(self, label_map: torch.Tensor, image: torch.Tensor | None = None) -> n_neutral = None # plain flip (sub-labels carry no L/R structure) randomise_bg = False # background is now modelled by its clusters - # 1. random crop to output_shape (label space) ------------------------ - if self.output_shape is not None and tuple(self.output_shape) != tuple(labels.shape[2:]): - labels = self._random_crop(labels, self.output_shape) - - # 2. spatial deformation of the LABEL MAP (nearest) ------------------- + # 1. spatial deformation of the LABEL MAP (nearest) ------------------- affine = None if self.apply_affine and self._affine_active(): affine = FN.sample_affine_matrices( @@ -283,6 +279,20 @@ def forward(self, label_map: torch.Tensor, image: torch.Tensor | None = None) -> .long() ) + # 2. random crop to output_shape (label space) ------------------------ + # + # After the deformation, not before. `warp_volume` pads with zeros, so + # cropping first meant the warp pulled background in from outside the crop + # box; lab2im's RandomSpatialDeformation(output_shape=...) resamples the + # full label map onto the output grid and therefore pulls in real anatomy. + # This is also the order the module docstring states. + # + # No shipped config is affected: `output_shape` defaults to None and is not + # among RandomSynthSegGPU's accepted parameters, so the crop only runs for a + # caller constructing SynthSegGenerator directly. + if self.output_shape is not None and tuple(self.output_shape) != tuple(labels.shape[2:]): + labels = self._random_crop(labels, self.output_shape) + # 3. left/right flipping (with optional label swap) ------------------- if self.flipping and float(torch.rand((), device=device)) < 0.5: label_values_flip = ( diff --git a/unit_tests/test_synthseg_crop_order.py b/unit_tests/test_synthseg_crop_order.py new file mode 100644 index 0000000..976d6e7 --- /dev/null +++ b/unit_tests/test_synthseg_crop_order.py @@ -0,0 +1,117 @@ +"""SynthSeg deforms the label map, then crops it -- not the other way round. + +The module docstring states the order as + + spatial deform (affine + diffeomorphic SVF, on labels, nearest) + -> [optional random crop] + -> left/right flip ... + +and lab2im's `RandomSpatialDeformation(output_shape=...)` resamples the *full* +label map onto the output grid. The code ran `_random_crop` first. + +That is not merely a stale comment. `warp_volume` pads with `zeros`, so +deforming a volume that has already been cropped pulls background in from +outside the crop box, where the reference implementation pulls in the real +anatomy that sits there. + +No shipped config is affected: `output_shape` defaults to None and is not among +`RandomSynthSegGPU`'s accepted parameters, so the crop only runs for a caller +constructing `SynthSegGenerator` directly. Both facts are asserted below so the +claim is checked rather than believed. +""" + +from __future__ import annotations + +import torch + +from smauglab import registry +from smauglab.registry import Backend +from smauglab.transforms.synthseg.generator import SynthSegGenerator +from unit_tests.helpers import SmaugLabTestCase + +FULL = (24, 24, 24) +CROPPED = (16, 16, 16) + + +def dense_labels(shape=FULL) -> torch.Tensor: + """A label map with no background at all, so any zero is imported padding.""" + labels = torch.ones(1, 1, *shape, dtype=torch.long) + labels[:, :, : shape[0] // 2] = 2 + return labels + + +class TestCropHappensAfterTheDeformation(SmaugLabTestCase): + def _generator(self, **kwargs): + return SynthSegGenerator( + generation_labels=[0, 1, 2], + output_labels=[0, 1, 2], + output_shape=CROPPED, + apply_affine=True, + apply_nonlinear=False, + flipping=False, + **kwargs, + ) + + def test_the_output_has_the_requested_shape(self): + torch.manual_seed(0) + + _, labels = self._generator()(dense_labels()) + + self.assertEqual(tuple(labels.shape[2:]), CROPPED) + + def test_far_less_background_is_imported(self): + """The behavioural difference, measured rather than asserted in prose. + + Every voxel of the input carries a label, so a background voxel in the + output can only be padding that `warp_volume` invented. Deforming an + already-cropped box pulls that in constantly, because the box is small and + the padding starts immediately outside it; deforming the full map first + only reaches padding when the crop lands against the volume edge. + + Measured over 40 seeds on a 24-cube input cropped to 16: + + crop -> deform: mean 13.5% background, 0/40 draws with none + deform -> crop: mean 3.3% background, 7/40 draws with none + + The thresholds sit between the two, so this fails on the old order. + """ + background = [] + for seed in range(40): + torch.manual_seed(seed) + + _, labels = self._generator()(dense_labels()) + + background.append(float((labels == 0).float().mean())) + + mean_background = sum(background) / len(background) + self.assertLess(mean_background, 0.08, f"too much padding imported: mean {mean_background:.4f}") + self.assertTrue( + any(value == 0.0 for value in background), + "no draw came through without imported padding, which the deform-first order should allow", + ) + + def test_the_order_is_the_one_the_docstring_states(self): + import inspect + + source = inspect.getsource(SynthSegGenerator.forward) + + self.assertLess( + source.index("spatial deformation of the LABEL MAP"), + source.index("random crop to output_shape"), + "the crop runs before the deformation again", + ) + + +class TestTheCropIsUnreachableFromAConfig(SmaugLabTestCase): + """Why no shipped pipeline changes.""" + + def test_output_shape_is_not_a_config_parameter(self): + registry.load_all() + entry = registry.get("RandomSynthSegGPU", Backend.GPU) + + self.assertNotIn("output_shape", registry.accepted_params(entry)) + + def test_the_generator_defaults_to_no_crop(self): + generator = SynthSegGenerator(generation_labels=[0, 1], output_labels=[0, 1]) + + self.assertIsNone(generator.output_shape)