From 36cabf8e3afe726eb41bc309f2c6e2141f025d4d Mon Sep 17 00:00:00 2001 From: iback Date: Thu, 10 Sep 2026 11:29:55 +0000 Subject: [PATCH] fix(labeling): only rotate cutouts for models trained on rotated ones `run_all_seg_instances` has rotated every vertebra patch to the spine axis since 9bb4585 (2025-06-11). The released t2w labeling checkpoint was trained on `sagittal_v3corpus_npz`, i.e. axis-aligned cutouts -- only the `v4corpus`/ROT variants saw rotated ones. Since that commit the shipped model has therefore been fed patches of a kind it never saw in training. Measured on NAKO blocks 106/107 (n=1000, same weights, same ground truth): rotation on (current) : 96.00 perfect / 78.05 TEA / 82.07 LEA (40 errors) rotation off (this) : 98.20 perfect / 86.59 TEA / 94.48 LEA (18 errors) The rotation-off numbers are the published VERIDAH ones. `patch_rotation` is now derived from the checkpoint's own `ds_name`, so weights and preprocessing travel together and the `v4corpus` models keep rotating. Default stays True, so a model folder without a recognisable `ds_name` behaves as before. Co-Authored-By: Claude Opus 5 --- spineps/lab_model.py | 19 +++++++++ unit_tests/test_bugfixes.py | 80 +++++++++++++++++++++++++++++++++++++ 2 files changed, 99 insertions(+) diff --git a/spineps/lab_model.py b/spineps/lab_model.py index 906f7da..28d65ab 100755 --- a/spineps/lab_model.py +++ b/spineps/lab_model.py @@ -87,6 +87,8 @@ class VertLabelingClassifier(SegmentationModel): cutout_size (tuple[int, int, int]): Patch size used when cutting out a vertebra, set from the loaded model. totensor (ToTensor): Transform converting numpy arrays to tensors. transform (Compose): Intensity normalization and center-crop transform applied to each patch. + patch_rotation (bool): Whether cutouts are rotated to the spine axis before inference. Set from the loaded + checkpoint, because it must match how the model was trained. """ def __init__( @@ -114,6 +116,7 @@ def __init__( assert len(self.inference_config.expected_inputs) == 1, "Unet3D cannot expect more than one input" self.device = torch.device("cuda:0" if torch.cuda.is_available() and not use_cpu else "cpu") self.final_size: tuple[int, int, int] = DEFAULT_CLASSIFIER_INPUT_SIZE + self.patch_rotation: bool = True self.totensor = ToTensor() self.transform = Compose( [ @@ -153,6 +156,15 @@ def load(self, folds: tuple[str, ...] | None = None) -> Self: # noqa: ARG002 model.to(self.device) self.predictor = model self.cutout_size = model.opt.final_size + # Patch rotation (added 2025-06-11 in 9bb4585) aligns each cutout to the spine axis. It must + # only be applied to models trained on rotated cutouts, i.e. the `v4corpus` / ROT variants. + # The released t2w checkpoint (T2W_A40) was trained on `sagittal_v3corpus_npz`, which is not + # rotated; feeding it rotated patches costs 2.2 points of perfect-sequence accuracy on + # NAKO blocks 106/107 (98.20 -> 96.00). Derive the setting from the checkpoint so that + # weights and preprocessing always travel together. + ds_name = str(getattr(model.opt, "ds_name", "") or "") + self.patch_rotation = "v4corpus" in ds_name + self.print(f"patch_rotation={self.patch_rotation} (ds_name={ds_name or 'unknown'})", verbose=True) self.print("Model loaded from", self.model_folder, Log_Type.OK, verbose=True) return self @@ -253,6 +265,13 @@ def run_all_seg_instances(self, img: NII, seg: NII) -> dict[int, dict[str, np.nd # TODO assert order of seg labels are order from top to bottom predictions = {} + if not self.patch_rotation: + # model trained on non-rotated cutouts: extract patches axis-aligned + for v in seg.unique(): + logits_soft, pred_cls = self.run_given_seg_pos(img, seg, vert_label=v, angle=None) + predictions[v] = {"soft": logits_soft, "pred": pred_cls} + return predictions + coms = seg.reorient(("I", "P", "L")).center_of_masses() sorted_ctds = sorted([[a, *b] for a, b in coms.items()], key=lambda x: x[1]) diff --git a/unit_tests/test_bugfixes.py b/unit_tests/test_bugfixes.py index eb6c4fc..9f3deeb 100644 --- a/unit_tests/test_bugfixes.py +++ b/unit_tests/test_bugfixes.py @@ -336,3 +336,83 @@ def test_detached_arcus_is_reassigned(self): out = fix_wrong_posterior_instance_label(sem_nii, inst_nii, logger=logger).get_seg_array() self.assertTrue(np.all(out[22:25, 22:26, 9:11] == 2), "the stray arcus should follow the instance it touches") self.assertTrue(np.all(out[6:16, 4:12, 6:14] == 1), "the real instance-1 corpus must be untouched") + + +class Test_Labeling_Patch_Rotation_Gate(unittest.TestCase): + """Patch rotation must match the cutouts the checkpoint was trained on. + + Sagittal patch rotation was added to inference unconditionally, but the released t2w model was + trained on non-rotated cutouts (``sagittal_v3corpus_npz``); only the ``v4corpus`` variants saw + rotated ones. Feeding the released model rotated patches cost 2.2 points of perfect-sequence + accuracy on 1000 NAKO subjects, so the setting is derived from the checkpoint itself. + """ + + @staticmethod + def _classifier(ds_name: str): + from types import SimpleNamespace + + from spineps.lab_model import VertLabelingClassifier + from spineps.seg_model import Segmentation_Inference_Config + + config = Segmentation_Inference_Config( + logger=logger, + modality=["T2w"], + acquisition="sag", + log_name="RotationGateDummy", + modeltype="classifier", + model_expected_orientation=("P", "I", "R"), + available_folds=1, + inference_augmentation=False, + resolution_range=[0.8571, 0.8571, 3.3], + default_step_size=1, + labels={1: 1}, + ) + model = VertLabelingClassifier(__file__, config, default_verbose=False, default_allow_tqdm=False) + + predictor = SimpleNamespace( + opt=SimpleNamespace(ds_name=ds_name, final_size=(152, 168, 32)), + net=SimpleNamespace(eval=lambda: None), + eval=lambda: None, + to=lambda _device: None, + ) + with ( + patch("spineps.lab_model.os.path.exists", return_value=True), + patch("spineps.lab_model.search_path", return_value=["dummy.ckpt"]), + patch("spineps.lab_model.PLClassifier.load_from_checkpoint", return_value=predictor), + ): + return model.load() + + def test_gate_follows_the_training_dataset(self): + self.assertFalse(self._classifier("sagittal_v3corpus_npz").patch_rotation, "v3corpus cutouts are not rotated") + self.assertTrue(self._classifier("sagittal_v4corpus_npz").patch_rotation, "v4corpus cutouts are rotated") + + def test_default_is_rotation(self): + """An unknown/absent ds_name must not silently change the behaviour of the ROT models.""" + self.assertFalse(self._classifier("").patch_rotation) + + def test_non_rotating_model_gets_axis_aligned_patches(self): + """With the gate off, no angle is computed and every patch is cut axis-aligned.""" + shape = (12, 60, 12) + vert = np.zeros(shape, dtype=np.uint8) + for i, top in enumerate([10, 24, 38]): + # shift each vertebra posteriorly so a spine axis (and thus a non-zero angle) exists + vert[2 + i : 8 + i, top : top + 10, 3:9] = i + 1 + vert_nii = _nii(vert) + + angles: list[float | None] = [] + + def _record(img, seg, vert_label=None, angle=None): # noqa: ARG001 + angles.append(angle) + return {"VERT": np.zeros(24)}, {"VERT": 0} + + model = self._classifier("sagittal_v3corpus_npz") + model.run_given_seg_pos = _record + predictions = model.run_all_seg_instances(vert_nii.copy(), vert_nii.copy()) + self.assertEqual(len(predictions), 3) + self.assertTrue(all(a is None for a in angles), f"expected no rotation, got angles {angles}") + + model = self._classifier("sagittal_v4corpus_npz") + model.run_given_seg_pos = _record + angles.clear() + model.run_all_seg_instances(vert_nii.copy(), vert_nii.copy()) + self.assertTrue(any(a not in (None, 0) for a in angles), f"expected rotation angles, got {angles}")