diff --git a/CMakeLists.txt b/CMakeLists.txt index e12d235..cdfa2c4 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -14,6 +14,7 @@ nanobind_add_module(_core src/bindings/blocking.cxx src/bindings/module.cxx src/bindings/graph.cxx + src/bindings/ground_truth.cxx src/bindings/segmentation.cxx src/bindings/utils.cxx src/cpp/segmentation/mutex_watershed.cxx diff --git a/MIGRATION_GUIDE.md b/MIGRATION_GUIDE.md index 22a4591..bdfe1dc 100644 --- a/MIGRATION_GUIDE.md +++ b/MIGRATION_GUIDE.md @@ -16,7 +16,8 @@ import bioimage_cpp as bic ``` Graph functionality is under `bic.graph`, segmentation functionality is under -`bic.segmentation`, and small utility functions are under `bic.utils`. +`bic.segmentation`, ground-truth comparison functionality is under +`bic.ground_truth`, and small utility functions are under `bic.utils`. ## Blocking @@ -220,6 +221,74 @@ Notes: - `number_of_threads=0` uses the library default; pass a positive integer for a fixed thread count. +## Segmentation Overlaps + +Nifty: + +```python +import nifty.ground_truth as ngt + +overlap = ngt.overlap(segmentation, ground_truth) +``` + +bioimage-cpp: + +```python +import bioimage_cpp as bic + +overlap = bic.ground_truth.segmentation_overlap(segmentation, ground_truth) +``` + +The first input is called `labels_a` and the second input is called `labels_b` +in the `bioimage-cpp` API. Use named structured tables instead of positional +arrays: + +```python +table = overlap.overlap_table() +# fields: "label_a", "label_b", "count" + +table = overlap.overlap_table(normalize_by="a") +# fields: "label_a", "label_b", "count", "fraction" + +overlaps = overlap.overlaps_for_label_a(12, normalize=True) +# fields: "label", "count", "fraction" + +best = overlap.best_overlap_for_label_a(12, ignore_zero=True) +# BestOverlap(label=..., count=..., fraction=..., found=...) +``` + +Other common queries: + +```python +overlap.labels_a +overlap.labels_b +overlap.count_a(12) +overlap.count_b(4) +overlap.overlap_count(12, 4) +overlap.counts_a_table() +overlap.counts_b_table() +overlap.best_overlap_for_label_b(4) +overlap.is_label_a_overlapping_with_zero(12) +overlap.different_overlap(12, 13) +``` + +Intentional improvements over nifty: + +- Labels are stored sparsely, so large sparse label ids do not require a dense + vector up to `max_label + 1`. +- The Python API returns structured arrays with named fields and a + `BestOverlap` dataclass instead of ambiguous positional arrays. +- Both overlap directions are supported explicitly: + `overlaps_for_label_a(...)` and `overlaps_for_label_b(...)`. +- Normalization is explicit via `normalize_by="a"`, `"b"`, or `"total"`. +- Missing labels return count `0`; best-overlap queries expose `found=False`. + +Notes: + +- Inputs must be integer arrays with identical shape. +- Signed integer inputs must not contain negative labels. +- Inputs are converted to contiguous `uint64` arrays before entering C++. + ## Mutex Watershed Affogato: diff --git a/include/bioimage_cpp/overlap.hxx b/include/bioimage_cpp/overlap.hxx new file mode 100644 index 0000000..6ac9a58 --- /dev/null +++ b/include/bioimage_cpp/overlap.hxx @@ -0,0 +1,262 @@ +#pragma once + +#include "bioimage_cpp/array_view.hxx" + +#include +#include +#include +#include +#include +#include +#include + +namespace bioimage_cpp::ground_truth { + +struct OverlapPair { + std::uint64_t label_a = 0; + std::uint64_t label_b = 0; + std::uint64_t count = 0; +}; + +class SegmentationOverlap { +public: + void add(const std::uint64_t label_a, const std::uint64_t label_b) { + ++overlaps_a_[label_a][label_b]; + ++overlaps_b_[label_b][label_a]; + ++counts_a_[label_a]; + ++counts_b_[label_b]; + ++total_count_; + } + + std::uint64_t total_count() const { + return total_count_; + } + + std::uint64_t count_a(const std::uint64_t label) const { + return map_value_or_zero(counts_a_, label); + } + + std::uint64_t count_b(const std::uint64_t label) const { + return map_value_or_zero(counts_b_, label); + } + + std::uint64_t overlap_count( + const std::uint64_t label_a, + const std::uint64_t label_b + ) const { + const auto found_a = overlaps_a_.find(label_a); + if (found_a == overlaps_a_.end()) { + return 0; + } + return map_value_or_zero(found_a->second, label_b); + } + + std::vector labels_a() const { + return sorted_keys(counts_a_); + } + + std::vector labels_b() const { + return sorted_keys(counts_b_); + } + + std::vector> counts_a() const { + return sorted_label_counts(counts_a_); + } + + std::vector> counts_b() const { + return sorted_label_counts(counts_b_); + } + + std::vector overlap_pairs() const { + std::vector result; + for (const auto &label_overlaps : overlaps_a_) { + for (const auto &overlap : label_overlaps.second) { + result.push_back(OverlapPair{ + label_overlaps.first, + overlap.first, + overlap.second, + }); + } + } + sort_overlap_pairs(result); + return result; + } + + std::vector> overlaps_for_label_a( + const std::uint64_t label_a + ) const { + return overlaps_for_label(overlaps_a_, label_a); + } + + std::vector> overlaps_for_label_b( + const std::uint64_t label_b + ) const { + return overlaps_for_label(overlaps_b_, label_b); + } + + std::pair best_overlap_for_label_a( + const std::uint64_t label_a, + const bool ignore_zero = false + ) const { + return best_overlap(overlaps_a_, label_a, ignore_zero); + } + + std::pair best_overlap_for_label_b( + const std::uint64_t label_b, + const bool ignore_zero = false + ) const { + return best_overlap(overlaps_b_, label_b, ignore_zero); + } + + bool is_label_a_overlapping_with_zero(const std::uint64_t label_a) const { + return overlap_count(label_a, 0) != 0; + } + + bool is_label_b_overlapping_with_zero(const std::uint64_t label_b) const { + const auto found_b = overlaps_b_.find(label_b); + if (found_b == overlaps_b_.end()) { + return false; + } + return found_b->second.find(0) != found_b->second.end(); + } + + double different_overlap(const std::uint64_t label_a_u, const std::uint64_t label_a_v) const { + const auto found_u = overlaps_a_.find(label_a_u); + const auto found_v = overlaps_a_.find(label_a_v); + if (found_u == overlaps_a_.end() || found_v == overlaps_a_.end()) { + throw std::out_of_range("labels must exist in segmentation A"); + } + + const auto size_u = static_cast(count_a(label_a_u)); + const auto size_v = static_cast(count_a(label_a_v)); + double result = 0.0; + for (const auto &overlap_u : found_u->second) { + for (const auto &overlap_v : found_v->second) { + if (overlap_u.first != overlap_v.first) { + result += + (static_cast(overlap_u.second) / size_u) * + (static_cast(overlap_v.second) / size_v); + } + } + } + return result; + } + +private: + using CountMap = std::unordered_map; + using OverlapMap = std::unordered_map; + + static std::uint64_t map_value_or_zero( + const CountMap &map, + const std::uint64_t label + ) { + const auto found = map.find(label); + return found == map.end() ? 0 : found->second; + } + + static std::vector sorted_keys(const CountMap &map) { + std::vector result; + result.reserve(map.size()); + for (const auto &entry : map) { + result.push_back(entry.first); + } + std::sort(result.begin(), result.end()); + return result; + } + + static std::vector> sorted_label_counts( + const CountMap &map + ) { + std::vector> result; + result.reserve(map.size()); + for (const auto &entry : map) { + result.push_back(entry); + } + std::sort(result.begin(), result.end(), [](const auto &a, const auto &b) { + return a.first < b.first; + }); + return result; + } + + static void sort_overlap_pairs(std::vector &pairs) { + std::sort(pairs.begin(), pairs.end(), [](const auto &a, const auto &b) { + if (a.label_a != b.label_a) { + return a.label_a < b.label_a; + } + return a.label_b < b.label_b; + }); + } + + static std::vector> overlaps_for_label( + const OverlapMap &overlaps, + const std::uint64_t label + ) { + const auto found = overlaps.find(label); + if (found == overlaps.end()) { + return {}; + } + return sorted_label_counts(found->second); + } + + static std::pair best_overlap( + const OverlapMap &overlaps, + const std::uint64_t label, + const bool ignore_zero + ) { + const auto found = overlaps.find(label); + if (found == overlaps.end()) { + return {0, 0}; + } + + std::uint64_t best_label = 0; + std::uint64_t best_count = 0; + for (const auto &overlap : found->second) { + if (ignore_zero && overlap.first == 0) { + continue; + } + if ( + overlap.second > best_count || + (overlap.second == best_count && overlap.first < best_label) + ) { + best_label = overlap.first; + best_count = overlap.second; + } + } + return {best_label, best_count}; + } + + CountMap counts_a_; + CountMap counts_b_; + OverlapMap overlaps_a_; + OverlapMap overlaps_b_; + std::uint64_t total_count_ = 0; +}; + +inline std::uint64_t array_size(const std::vector &shape) { + std::uint64_t size = 1; + for (const auto axis_size : shape) { + if (axis_size < 0) { + throw std::invalid_argument("shape entries must be non-negative"); + } + size *= static_cast(axis_size); + } + return size; +} + +inline SegmentationOverlap segmentation_overlap( + const ConstArrayView &labels_a, + const ConstArrayView &labels_b +) { + if (labels_a.shape != labels_b.shape) { + throw std::invalid_argument("labels_a and labels_b must have the same shape"); + } + + SegmentationOverlap result; + const auto size = array_size(labels_a.shape); + for (std::uint64_t index = 0; index < size; ++index) { + result.add(labels_a.data[index], labels_b.data[index]); + } + return result; +} + +} // namespace bioimage_cpp::ground_truth diff --git a/src/bindings/ground_truth.cxx b/src/bindings/ground_truth.cxx new file mode 100644 index 0000000..3d454aa --- /dev/null +++ b/src/bindings/ground_truth.cxx @@ -0,0 +1,117 @@ +#include "ground_truth.hxx" + +#include "bioimage_cpp/array_view.hxx" +#include "bioimage_cpp/overlap.hxx" + +#include +#include +#include + +#include +#include +#include + +namespace nb = nanobind; + +namespace bioimage_cpp::bindings { +namespace { + +using LabelArray = nb::ndarray; +using SegmentationOverlap = ground_truth::SegmentationOverlap; + +std::vector ndarray_shape(LabelArray array) { + std::vector shape(array.ndim()); + for (std::size_t axis = 0; axis < array.ndim(); ++axis) { + shape[axis] = static_cast(array.shape(axis)); + } + return shape; +} + +SegmentationOverlap segmentation_overlap(LabelArray labels_a, LabelArray labels_b) { + ConstArrayView labels_a_view{ + labels_a.data(), + ndarray_shape(labels_a), + {}, + }; + ConstArrayView labels_b_view{ + labels_b.data(), + ndarray_shape(labels_b), + {}, + }; + + nb::gil_scoped_release release; + return ground_truth::segmentation_overlap(labels_a_view, labels_b_view); +} + +} // namespace + +void bind_ground_truth(nb::module_ &m) { + nb::class_(m, "_OverlapPair") + .def_ro("label_a", &ground_truth::OverlapPair::label_a) + .def_ro("label_b", &ground_truth::OverlapPair::label_b) + .def_ro("count", &ground_truth::OverlapPair::count); + + nb::class_(m, "_SegmentationOverlap") + .def_prop_ro("total_count", &SegmentationOverlap::total_count) + .def("labels_a", &SegmentationOverlap::labels_a) + .def("labels_b", &SegmentationOverlap::labels_b) + .def("counts_a", &SegmentationOverlap::counts_a) + .def("counts_b", &SegmentationOverlap::counts_b) + .def("count_a", &SegmentationOverlap::count_a, nb::arg("label")) + .def("count_b", &SegmentationOverlap::count_b, nb::arg("label")) + .def( + "overlap_count", + &SegmentationOverlap::overlap_count, + nb::arg("label_a"), + nb::arg("label_b") + ) + .def("overlap_pairs", &SegmentationOverlap::overlap_pairs) + .def( + "overlaps_for_label_a", + &SegmentationOverlap::overlaps_for_label_a, + nb::arg("label") + ) + .def( + "overlaps_for_label_b", + &SegmentationOverlap::overlaps_for_label_b, + nb::arg("label") + ) + .def( + "best_overlap_for_label_a", + &SegmentationOverlap::best_overlap_for_label_a, + nb::arg("label"), + nb::arg("ignore_zero") = false + ) + .def( + "best_overlap_for_label_b", + &SegmentationOverlap::best_overlap_for_label_b, + nb::arg("label"), + nb::arg("ignore_zero") = false + ) + .def( + "is_label_a_overlapping_with_zero", + &SegmentationOverlap::is_label_a_overlapping_with_zero, + nb::arg("label") + ) + .def( + "is_label_b_overlapping_with_zero", + &SegmentationOverlap::is_label_b_overlapping_with_zero, + nb::arg("label") + ) + .def( + "different_overlap", + &SegmentationOverlap::different_overlap, + nb::arg("label_a_u"), + nb::arg("label_a_v") + ); + + m.def( + "_segmentation_overlap_uint64", + &segmentation_overlap, + nb::arg("labels_a"), + nb::arg("labels_b"), + "Compute sparse overlap counts between two uint64 label arrays." + ); +} + +} // namespace bioimage_cpp::bindings diff --git a/src/bindings/ground_truth.hxx b/src/bindings/ground_truth.hxx new file mode 100644 index 0000000..51294e3 --- /dev/null +++ b/src/bindings/ground_truth.hxx @@ -0,0 +1,9 @@ +#pragma once + +#include + +namespace bioimage_cpp::bindings { + +void bind_ground_truth(nanobind::module_ &m); + +} // namespace bioimage_cpp::bindings diff --git a/src/bindings/module.cxx b/src/bindings/module.cxx index 8b9206a..cee23a9 100644 --- a/src/bindings/module.cxx +++ b/src/bindings/module.cxx @@ -1,5 +1,6 @@ #include "blocking.hxx" #include "graph.hxx" +#include "ground_truth.hxx" #include "segmentation.hxx" #include "utils.hxx" @@ -11,6 +12,7 @@ NB_MODULE(_core, m) { m.doc() = "C++ extension module for bioimage_cpp."; bioimage_cpp::bindings::bind_blocking(m); bioimage_cpp::bindings::bind_graph(m); + bioimage_cpp::bindings::bind_ground_truth(m); bioimage_cpp::bindings::bind_segmentation(m); bioimage_cpp::bindings::bind_utils(m); } diff --git a/src/bioimage_cpp/__init__.py b/src/bioimage_cpp/__init__.py index 4879fb3..f8a5e6f 100644 --- a/src/bioimage_cpp/__init__.py +++ b/src/bioimage_cpp/__init__.py @@ -3,6 +3,7 @@ from ._version import __version__ from ._core import Block, Blocking, BlockWithHalo from . import graph +from . import ground_truth from . import segmentation from . import utils @@ -12,6 +13,7 @@ "Blocking", "BlockWithHalo", "graph", + "ground_truth", "segmentation", "utils", ] diff --git a/src/bioimage_cpp/ground_truth/__init__.py b/src/bioimage_cpp/ground_truth/__init__.py new file mode 100644 index 0000000..00c6a93 --- /dev/null +++ b/src/bioimage_cpp/ground_truth/__init__.py @@ -0,0 +1,298 @@ +"""Ground-truth comparison helpers.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + +import numpy as np + +from .. import _core + +_COUNT_TABLE_DTYPE = np.dtype([("label", np.uint64), ("count", np.uint64)]) +_OVERLAP_TABLE_DTYPE = np.dtype( + [("label_a", np.uint64), ("label_b", np.uint64), ("count", np.uint64)] +) +_OVERLAP_FRACTION_TABLE_DTYPE = np.dtype( + [ + ("label_a", np.uint64), + ("label_b", np.uint64), + ("count", np.uint64), + ("fraction", np.float64), + ] +) +_LABEL_OVERLAP_TABLE_DTYPE = np.dtype( + [("label", np.uint64), ("count", np.uint64)] +) +_LABEL_OVERLAP_FRACTION_TABLE_DTYPE = np.dtype( + [("label", np.uint64), ("count", np.uint64), ("fraction", np.float64)] +) + + +@dataclass(frozen=True) +class BestOverlap: + """Best overlap result for one queried label.""" + + label: int + count: int + fraction: float + found: bool + + +class SegmentationOverlap: + """Sparse overlap counts between two segmentation label arrays. + + Use :func:`segmentation_overlap` to construct this object. Labels from the + first input are called ``label_a`` and labels from the second input are + called ``label_b`` in all tables. + """ + + def __init__(self, core_overlap): + self._core_overlap = core_overlap + + @property + def total_count(self) -> int: + """Total number of pixels or voxels compared.""" + return int(self._core_overlap.total_count) + + @property + def labels_a(self) -> np.ndarray: + """Sorted labels present in the first segmentation.""" + return np.asarray(self._core_overlap.labels_a(), dtype=np.uint64) + + @property + def labels_b(self) -> np.ndarray: + """Sorted labels present in the second segmentation.""" + return np.asarray(self._core_overlap.labels_b(), dtype=np.uint64) + + def count_a(self, label: int) -> int: + """Return the size of a label in the first segmentation.""" + return int(self._core_overlap.count_a(_normalize_label(label))) + + def count_b(self, label: int) -> int: + """Return the size of a label in the second segmentation.""" + return int(self._core_overlap.count_b(_normalize_label(label))) + + def overlap_count(self, label_a: int, label_b: int) -> int: + """Return the number of pixels where ``label_a`` and ``label_b`` overlap.""" + return int( + self._core_overlap.overlap_count( + _normalize_label(label_a, "label_a"), + _normalize_label(label_b, "label_b"), + ) + ) + + def counts_a_table(self) -> np.ndarray: + """Return a structured array with fields ``label`` and ``count``.""" + return _label_count_table(self._core_overlap.counts_a()) + + def counts_b_table(self) -> np.ndarray: + """Return a structured array with fields ``label`` and ``count``.""" + return _label_count_table(self._core_overlap.counts_b()) + + def overlap_table( + self, + *, + normalize_by: Literal["a", "b", "total"] | None = None, + ) -> np.ndarray: + """Return all non-zero overlaps as a structured array. + + Without normalization, the fields are ``label_a``, ``label_b`` and + ``count``. With ``normalize_by="a"``, ``"b"`` or ``"total"``, a + ``fraction`` field is added. + """ + rows = self._core_overlap.overlap_pairs() + if normalize_by is None: + table = np.empty(len(rows), dtype=_OVERLAP_TABLE_DTYPE) + for index, row in enumerate(rows): + table[index] = (row.label_a, row.label_b, row.count) + return table + + _validate_normalize_by(normalize_by) + table = np.empty(len(rows), dtype=_OVERLAP_FRACTION_TABLE_DTYPE) + for index, row in enumerate(rows): + table[index] = ( + row.label_a, + row.label_b, + row.count, + self._fraction(row.label_a, row.label_b, row.count, normalize_by), + ) + return table + + def overlaps_for_label_a( + self, + label: int, + *, + normalize: bool = False, + ) -> np.ndarray: + """Return labels from B overlapping one label from A.""" + normalized_label = _normalize_label(label) + rows = self._core_overlap.overlaps_for_label_a(normalized_label) + denominator = self.count_a(label) + return _label_overlap_table(rows, normalize=normalize, denominator=denominator) + + def overlaps_for_label_b( + self, + label: int, + *, + normalize: bool = False, + ) -> np.ndarray: + """Return labels from A overlapping one label from B.""" + normalized_label = _normalize_label(label) + rows = self._core_overlap.overlaps_for_label_b(normalized_label) + denominator = self.count_b(label) + return _label_overlap_table(rows, normalize=normalize, denominator=denominator) + + def best_overlap_for_label_a( + self, + label: int, + *, + ignore_zero: bool = False, + ) -> BestOverlap: + """Return the best matching label in B for one label in A.""" + normalized_label = _normalize_label(label) + best_label, count = self._core_overlap.best_overlap_for_label_a( + normalized_label, bool(ignore_zero) + ) + denominator = self.count_a(label) + return _best_overlap(best_label, count, denominator) + + def best_overlap_for_label_b( + self, + label: int, + *, + ignore_zero: bool = False, + ) -> BestOverlap: + """Return the best matching label in A for one label in B.""" + normalized_label = _normalize_label(label) + best_label, count = self._core_overlap.best_overlap_for_label_b( + normalized_label, bool(ignore_zero) + ) + denominator = self.count_b(label) + return _best_overlap(best_label, count, denominator) + + def is_label_a_overlapping_with_zero(self, label: int) -> bool: + """Return whether a label from A overlaps label ``0`` in B.""" + return bool( + self._core_overlap.is_label_a_overlapping_with_zero(_normalize_label(label)) + ) + + def is_label_b_overlapping_with_zero(self, label: int) -> bool: + """Return whether a label from B overlaps label ``0`` in A.""" + return bool( + self._core_overlap.is_label_b_overlapping_with_zero(_normalize_label(label)) + ) + + def different_overlap(self, label_a_u: int, label_a_v: int) -> float: + """Return the probability that two A labels overlap different B labels.""" + return float( + self._core_overlap.different_overlap( + _normalize_label(label_a_u, "label_a_u"), + _normalize_label(label_a_v, "label_a_v"), + ) + ) + + def _fraction( + self, + label_a: int, + label_b: int, + count: int, + normalize_by: str, + ) -> float: + if normalize_by == "a": + denominator = self.count_a(label_a) + elif normalize_by == "b": + denominator = self.count_b(label_b) + else: + denominator = self.total_count + return 0.0 if denominator == 0 else float(count) / float(denominator) + + +def segmentation_overlap(labels_a: np.ndarray, labels_b: np.ndarray) -> SegmentationOverlap: + """Compute sparse overlap counts between two segmentations. + + Parameters + ---------- + labels_a, labels_b: + Integer NumPy arrays with identical shape. Supported dtypes are + unsigned integer dtypes and signed integer dtypes with non-negative + values. Inputs are converted to contiguous ``uint64`` arrays before + entering C++. + + Returns + ------- + SegmentationOverlap + Object exposing named tables and query methods for overlap counts. + """ + array_a = _normalize_labels(labels_a, "labels_a") + array_b = _normalize_labels(labels_b, "labels_b") + if array_a.shape != array_b.shape: + raise ValueError( + "labels_a and labels_b must have the same shape, got " + f"labels_a shape={array_a.shape}, labels_b shape={array_b.shape}" + ) + if array_a.ndim == 0: + raise ValueError("labels_a and labels_b must have at least one dimension") + + return SegmentationOverlap(_core._segmentation_overlap_uint64(array_a, array_b)) + + +def _normalize_labels(labels: np.ndarray, name: str) -> np.ndarray: + array = np.asarray(labels) + if not np.issubdtype(array.dtype, np.integer): + raise TypeError(f"{name} must have an integer dtype, got dtype={array.dtype}") + if array.ndim == 0: + raise ValueError(f"{name} must have at least one dimension") + if np.issubdtype(array.dtype, np.signedinteger) and np.any(array < 0): + raise ValueError(f"{name} must not contain negative labels") + return np.ascontiguousarray(array, dtype=np.uint64) + + +def _normalize_label(label: int, name: str = "label") -> int: + label = int(label) + if label < 0: + raise ValueError(f"{name} must be non-negative") + return label + + +def _label_count_table(rows) -> np.ndarray: + table = np.empty(len(rows), dtype=_COUNT_TABLE_DTYPE) + for index, (label, count) in enumerate(rows): + table[index] = (label, count) + return table + + +def _label_overlap_table(rows, *, normalize: bool, denominator: int) -> np.ndarray: + if not normalize: + table = np.empty(len(rows), dtype=_LABEL_OVERLAP_TABLE_DTYPE) + for index, (label, count) in enumerate(rows): + table[index] = (label, count) + return table + + table = np.empty(len(rows), dtype=_LABEL_OVERLAP_FRACTION_TABLE_DTYPE) + for index, (label, count) in enumerate(rows): + fraction = 0.0 if denominator == 0 else float(count) / float(denominator) + table[index] = (label, count, fraction) + return table + + +def _best_overlap(label: int, count: int, denominator: int) -> BestOverlap: + fraction = 0.0 if denominator == 0 else float(count) / float(denominator) + return BestOverlap( + label=int(label), + count=int(count), + fraction=fraction, + found=bool(count), + ) + + +def _validate_normalize_by(normalize_by: str) -> None: + if normalize_by not in ("a", "b", "total"): + raise ValueError("normalize_by must be one of None, 'a', 'b', or 'total'") + + +__all__ = [ + "BestOverlap", + "SegmentationOverlap", + "segmentation_overlap", +] diff --git a/tests/test_ground_truth_overlap.py b/tests/test_ground_truth_overlap.py new file mode 100644 index 0000000..55f0f2b --- /dev/null +++ b/tests/test_ground_truth_overlap.py @@ -0,0 +1,139 @@ +import numpy as np +import pytest + +import bioimage_cpp as bic + + +def test_segmentation_overlap_tables_and_counts(): + labels_a = np.array([[1, 1, 2], [1, 3, 2]], dtype=np.uint64) + labels_b = np.array([[5, 5, 5], [6, 6, 7]], dtype=np.uint32) + + overlap = bic.ground_truth.segmentation_overlap(labels_a, labels_b) + + assert overlap.total_count == 6 + np.testing.assert_array_equal(overlap.labels_a, np.array([1, 2, 3], dtype=np.uint64)) + np.testing.assert_array_equal(overlap.labels_b, np.array([5, 6, 7], dtype=np.uint64)) + assert overlap.count_a(1) == 3 + assert overlap.count_b(5) == 3 + assert overlap.overlap_count(1, 5) == 2 + assert overlap.overlap_count(42, 5) == 0 + + counts_a = overlap.counts_a_table() + assert counts_a.dtype.names == ("label", "count") + np.testing.assert_array_equal(counts_a["label"], [1, 2, 3]) + np.testing.assert_array_equal(counts_a["count"], [3, 2, 1]) + + table = overlap.overlap_table() + assert table.dtype.names == ("label_a", "label_b", "count") + np.testing.assert_array_equal(table["label_a"], [1, 1, 2, 2, 3]) + np.testing.assert_array_equal(table["label_b"], [5, 6, 5, 7, 6]) + np.testing.assert_array_equal(table["count"], [2, 1, 1, 1, 1]) + + +def test_segmentation_overlap_normalized_tables(): + labels_a = np.array([[1, 1, 2], [1, 3, 2]], dtype=np.uint64) + labels_b = np.array([[5, 5, 5], [6, 6, 7]], dtype=np.uint64) + overlap = bic.ground_truth.segmentation_overlap(labels_a, labels_b) + + by_a = overlap.overlap_table(normalize_by="a") + assert by_a.dtype.names == ("label_a", "label_b", "count", "fraction") + np.testing.assert_allclose(by_a["fraction"], [2 / 3, 1 / 3, 1 / 2, 1 / 2, 1.0]) + + by_b = overlap.overlap_table(normalize_by="b") + np.testing.assert_allclose(by_b["fraction"], [2 / 3, 1 / 2, 1 / 3, 1.0, 1 / 2]) + + by_total = overlap.overlap_table(normalize_by="total") + np.testing.assert_allclose(by_total["fraction"], [2 / 6, 1 / 6, 1 / 6, 1 / 6, 1 / 6]) + + with pytest.raises(ValueError, match="normalize_by"): + overlap.overlap_table(normalize_by="bad") + + +def test_per_label_overlaps_and_best_overlap_are_clear(): + labels_a = np.array([[1, 1, 2], [1, 3, 2]], dtype=np.uint64) + labels_b = np.array([[5, 5, 5], [6, 6, 7]], dtype=np.uint64) + overlap = bic.ground_truth.segmentation_overlap(labels_a, labels_b) + + overlaps_a = overlap.overlaps_for_label_a(1, normalize=True) + assert overlaps_a.dtype.names == ("label", "count", "fraction") + np.testing.assert_array_equal(overlaps_a["label"], [5, 6]) + np.testing.assert_array_equal(overlaps_a["count"], [2, 1]) + np.testing.assert_allclose(overlaps_a["fraction"], [2 / 3, 1 / 3]) + + overlaps_b = overlap.overlaps_for_label_b(5) + assert overlaps_b.dtype.names == ("label", "count") + np.testing.assert_array_equal(overlaps_b["label"], [1, 2]) + np.testing.assert_array_equal(overlaps_b["count"], [2, 1]) + + best_a = overlap.best_overlap_for_label_a(1) + assert best_a.label == 5 + assert best_a.count == 2 + assert best_a.fraction == pytest.approx(2 / 3) + assert best_a.found + + best_b = overlap.best_overlap_for_label_b(5) + assert best_b.label == 1 + assert best_b.count == 2 + assert best_b.fraction == pytest.approx(2 / 3) + assert best_b.found + + missing = overlap.best_overlap_for_label_a(99) + assert missing == bic.ground_truth.BestOverlap(label=0, count=0, fraction=0.0, found=False) + + +def test_zero_label_handling_and_different_overlap(): + labels_a = np.array([1, 1, 2, 2, 3], dtype=np.uint64) + labels_b = np.array([0, 5, 0, 6, 6], dtype=np.uint64) + overlap = bic.ground_truth.segmentation_overlap(labels_a, labels_b) + + assert overlap.is_label_a_overlapping_with_zero(1) + assert overlap.is_label_b_overlapping_with_zero(0) is False + assert overlap.best_overlap_for_label_a(1).label == 0 + assert overlap.best_overlap_for_label_a(1, ignore_zero=True).label == 5 + assert overlap.best_overlap_for_label_a(1, ignore_zero=True).found + assert not overlap.best_overlap_for_label_a(99, ignore_zero=True).found + + assert overlap.different_overlap(1, 2) == pytest.approx(0.75) + with pytest.raises(IndexError, match="labels must exist"): + overlap.different_overlap(1, 99) + with pytest.raises(ValueError, match="non-negative"): + overlap.count_a(-1) + + +def test_sparse_large_labels_do_not_require_dense_max_label_storage(): + labels_a = np.array([1, 1_000_000_000_000], dtype=np.uint64) + labels_b = np.array([7, 8], dtype=np.uint64) + + overlap = bic.ground_truth.segmentation_overlap(labels_a, labels_b) + + np.testing.assert_array_equal( + overlap.labels_a, + np.array([1, 1_000_000_000_000], dtype=np.uint64), + ) + np.testing.assert_array_equal(overlap.overlap_table()["count"], [1, 1]) + + +def test_overlap_rejects_invalid_inputs(): + with pytest.raises(ValueError, match="same shape"): + bic.ground_truth.segmentation_overlap( + np.ones((2, 2), dtype=np.uint64), + np.ones((2, 3), dtype=np.uint64), + ) + + with pytest.raises(TypeError, match="integer dtype"): + bic.ground_truth.segmentation_overlap( + np.ones((2,), dtype=np.float32), + np.ones((2,), dtype=np.uint64), + ) + + with pytest.raises(ValueError, match="negative labels"): + bic.ground_truth.segmentation_overlap( + np.array([-1, 2], dtype=np.int64), + np.ones((2,), dtype=np.uint64), + ) + + with pytest.raises(ValueError, match="at least one dimension"): + bic.ground_truth.segmentation_overlap( + np.array(1, dtype=np.uint64), + np.array(1, dtype=np.uint64), + )