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
42 changes: 39 additions & 3 deletions development/distance/benchmark.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,8 @@
outside the timing loop so per-call dtype conversion does not show up in the
measurement of a particular library:

* bioimage_cpp.distance.distance_transform — uint8 mask
* bioimage_cpp.distance.distance_transform — uint8 mask
* bioimage_cpp.distance.distance_transform + indices — uint8 mask
* bioimage_cpp.distance.vector_difference_transform — uint8 mask
* vigra.filters.distanceTransform / vectorDistanceTransform — float32 mask
* scipy.ndimage.distance_transform_edt — float32 mask
Expand Down Expand Up @@ -43,7 +44,11 @@


LIBRARIES = ("bioimage_cpp", "vigra", "scipy")
OPERATIONS = ("distance_transform", "vector_difference_transform")
OPERATIONS = (
"distance_transform",
"distance_transform_indices",
"vector_difference_transform",
)


@dataclass(frozen=True)
Expand Down Expand Up @@ -246,6 +251,21 @@ def fn(mask: np.ndarray) -> np.ndarray:
return fn


def _bic_distance_indices(sampling: tuple[float, ...], n_threads: int):
from bioimage_cpp import distance

def fn(mask: np.ndarray) -> np.ndarray:
distances, _ = distance.distance_transform(
mask,
sampling=sampling,
return_indices=True,
number_of_threads=n_threads,
)
return distances

return fn


def _vigra_distance(sampling: tuple[float, ...]):
_quiet_vigra_matplotlib_cache()
import vigra.filters as vf
Expand Down Expand Up @@ -300,6 +320,18 @@ def fn(mask: np.ndarray) -> np.ndarray:
return fn


def _scipy_distance_indices(sampling: tuple[float, ...]):
from scipy import ndimage

def fn(mask: np.ndarray) -> np.ndarray:
distances, _ = ndimage.distance_transform_edt(
mask, sampling=sampling, return_indices=True
)
return distances

return fn


def _prepare_mask(library: str, base_mask: np.ndarray) -> np.ndarray:
if library == "bioimage_cpp":
# The Python wrapper fast-paths uint8 C-contiguous input.
Expand All @@ -324,6 +356,10 @@ def build_adapters(
"vigra": lambda: _vigra_distance(sampling),
"scipy": lambda: _scipy_distance(sampling),
},
"distance_transform_indices": {
"bioimage_cpp": lambda: _bic_distance_indices(sampling, n_threads),
"scipy": lambda: _scipy_distance_indices(sampling),
},
"vector_difference_transform": {
"bioimage_cpp": lambda: _bic_vector(sampling, n_threads),
"vigra": lambda: _vigra_vector(sampling),
Expand Down Expand Up @@ -394,7 +430,7 @@ def check_results(
errors = {}
for library, adapter in adapters.items():
result = np.asarray(adapter.fn(adapter.mask))
if operation == "distance_transform":
if operation in ("distance_transform", "distance_transform_indices"):
errors[library] = float(
np.max(np.abs(result.astype(np.float32) - reference_distance))
)
Expand Down
Loading
Loading