Skip to content
Merged
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
19 changes: 19 additions & 0 deletions spineps/lab_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__(
Expand Down Expand Up @@ -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(
[
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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])

Expand Down
80 changes: 80 additions & 0 deletions unit_tests/test_bugfixes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Loading