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
48 changes: 43 additions & 5 deletions src/pyrecest/utils/_roi_assignment_otsu.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,35 @@
from ._roi_assignment_extreme_range import patch_similarity_assignment_extreme_range


def _contains_masked_value(value) -> bool:
"""Return whether *value* contains at least one masked NumPy entry."""

if np.ma.is_masked(value):
return True
if isinstance(value, np.ndarray):
if value.dtype != object:
return False
items = value.reshape(-1)
elif isinstance(value, (list, tuple)):
items = value
else:
return False
return any(_contains_masked_value(item) for item in items)


def _reject_masked_value(value, name: str) -> None:
"""Reject masks before backend conversion silently discards them."""

if _contains_masked_value(value):
raise ValueError(f"{name} must not contain masked values.")


def _as_positive_nbins(roi_assignment_module, nbins) -> int:
"""Reject temporal scalars before backend integer coercion loses their dtype."""
"""Reject temporal and masked scalars before backend integer coercion."""

message = "nbins must be a positive integer."
if _contains_masked_value(nbins):
raise ValueError(message)
try:
value_array = np.asarray(nbins)
except (TypeError, ValueError) as exc:
Expand All @@ -28,15 +53,23 @@ def _as_positive_nbins(roi_assignment_module, nbins) -> int:


def _patch_minimum_similarity_threshold(roi_assignment_module) -> None:
"""Validate histogram bin counts before minimum-threshold early returns."""
"""Validate histogram inputs before minimum-threshold early returns."""

original_minimum = roi_assignment_module.minimum_similarity_threshold
if getattr(original_minimum, "_pyrecest_positive_nbins_validation", False):
if (
getattr(original_minimum, "_pyrecest_positive_nbins_validation", False)
and getattr(
original_minimum,
"_pyrecest_masked_similarity_validation",
False,
)
):
return

def minimum_similarity_threshold(similarities, *, nbins: int = 256) -> float:
"""Estimate a threshold by locating a valley between the two strongest modes."""

_reject_masked_value(similarities, "similarities")
nbins = _as_positive_nbins(roi_assignment_module, nbins)
return original_minimum(similarities, nbins=nbins)

Expand All @@ -47,6 +80,7 @@ def minimum_similarity_threshold(similarities, *, nbins: int = 256) -> float:
)
minimum_similarity_threshold.__doc__ = getattr(original_minimum, "__doc__", None)
minimum_similarity_threshold._pyrecest_positive_nbins_validation = True
minimum_similarity_threshold._pyrecest_masked_similarity_validation = True
roi_assignment_module.minimum_similarity_threshold = minimum_similarity_threshold


Expand All @@ -55,15 +89,18 @@ def patch_otsu_similarity_threshold(roi_assignment_module) -> None:

patch_similarity_assignment_extreme_range(roi_assignment_module)
original_otsu = roi_assignment_module.otsu_similarity_threshold
if getattr(original_otsu, "_pyrecest_strict_foreground_split", False) and getattr(
original_otsu, "_pyrecest_positive_nbins_validation", False
if (
getattr(original_otsu, "_pyrecest_strict_foreground_split", False)
and getattr(original_otsu, "_pyrecest_positive_nbins_validation", False)
and getattr(original_otsu, "_pyrecest_masked_similarity_validation", False)
):
_patch_minimum_similarity_threshold(roi_assignment_module)
return

def otsu_similarity_threshold(similarities, *, nbins: int = 256) -> float:
"""Estimate a threshold using Otsu's method on one-dimensional similarities."""

_reject_masked_value(similarities, "similarities")
nbins = _as_positive_nbins(roi_assignment_module, nbins)
values = roi_assignment_module.asarray(
similarities,
Expand Down Expand Up @@ -140,5 +177,6 @@ def otsu_similarity_threshold(similarities, *, nbins: int = 256) -> float:
otsu_similarity_threshold.__doc__ = getattr(original_otsu, "__doc__", None)
otsu_similarity_threshold._pyrecest_strict_foreground_split = True
otsu_similarity_threshold._pyrecest_positive_nbins_validation = True
otsu_similarity_threshold._pyrecest_masked_similarity_validation = True
roi_assignment_module.otsu_similarity_threshold = otsu_similarity_threshold
_patch_minimum_similarity_threshold(roi_assignment_module)
58 changes: 58 additions & 0 deletions tests/test_roi_assignment_masked_threshold_inputs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
import unittest

import numpy as np
from pyrecest.utils.roi_assignment import (
minimum_similarity_threshold,
otsu_similarity_threshold,
)


class TestRoiAssignmentMaskedThresholdInputs(unittest.TestCase):
def test_thresholds_reject_masked_similarity_entries(self):
similarities = np.ma.array(
[0.05, 0.08, 0.81, 0.9],
mask=[False, True, False, False],
)

for threshold_fn in (otsu_similarity_threshold, minimum_similarity_threshold):
with self.subTest(threshold_fn=threshold_fn.__name__):
with self.assertRaisesRegex(
ValueError,
"similarities must not contain masked values",
):
threshold_fn(similarities)

def test_thresholds_reject_nested_masked_similarity_entries(self):
similarities = [0.05, np.ma.masked, 0.81, 0.9]

for threshold_fn in (otsu_similarity_threshold, minimum_similarity_threshold):
with self.subTest(threshold_fn=threshold_fn.__name__):
with self.assertRaisesRegex(
ValueError,
"similarities must not contain masked values",
):
threshold_fn(similarities)

def test_thresholds_reject_masked_bin_count(self):
similarities = np.array([0.05, 0.08, 0.81, 0.9])
masked_nbins = np.ma.array(32, mask=True)

for threshold_fn in (otsu_similarity_threshold, minimum_similarity_threshold):
with self.subTest(threshold_fn=threshold_fn.__name__):
with self.assertRaisesRegex(ValueError, "nbins"):
threshold_fn(similarities, nbins=masked_nbins)

def test_thresholds_accept_masked_array_without_masked_entries(self):
similarities = np.ma.array(
[0.05, 0.08, 0.81, 0.9],
mask=False,
)

for threshold_fn in (otsu_similarity_threshold, minimum_similarity_threshold):
with self.subTest(threshold_fn=threshold_fn.__name__):
threshold = threshold_fn(similarities, nbins=16)
self.assertTrue(np.isfinite(threshold))


if __name__ == "__main__":
unittest.main()
Loading