Skip to content

Commit 62b2cbb

Browse files
committed
feat(core): ✨ add dilate_labels post-segmentation wrapper
Grows label objects by N pixels after any segmentation method, applied per-tile before halo trim/merge so dilated labels still stitch correctly across tile boundaries. Method-agnostic: wraps any (tile) -> labels callable rather than being baked into a specific plugin.
1 parent 1cfd0e8 commit 62b2cbb

3 files changed

Lines changed: 144 additions & 0 deletions

File tree

src/patchworks/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@
3535
from ._distributed import create_stage, spatial_tiles, stage_tile
3636
from ._io import estimate_empty_tiles, load_ome_zarr
3737
from ._merge import merge_tile_labels
38+
from ._postprocess import dilate_labels
3839
from ._relabel import relabel_sequential_array, relabel_sequential_zarr
3940
from ._relations import label_relations
4041

@@ -57,4 +58,5 @@
5758
"spatial_tiles",
5859
"create_stage",
5960
"stage_tile",
61+
"dilate_labels",
6062
]

src/patchworks/_postprocess.py

Lines changed: 101 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,101 @@
1+
"""Generic post-segmentation wrappers for patchworks.
2+
3+
These wrap any ``fn(tile) -> labels`` callable (a plugin, a custom function,
4+
whatever ``method`` in the Snakemake workflow builds) so the same
5+
post-processing applies regardless of which segmentation method produced the
6+
labels.
7+
8+
Usage
9+
-----
10+
>>> from patchworks import tile_process, dilate_labels
11+
>>> from patchworks.plugins.dog import dog_label_fn
12+
>>>
13+
>>> fn = dog_label_fn(low_sigma=1.0, high_sigma=3.0, threshold=0.02)
14+
>>> fn = dilate_labels(fn, iterations=2)
15+
>>> result = tile_process("image.zarr", fn, tile_shape=(1, 2048, 2048),
16+
... overlap=8, write_to="labels.zarr")
17+
"""
18+
19+
from __future__ import annotations
20+
21+
from functools import partial
22+
from typing import Callable
23+
24+
import numpy as np
25+
26+
27+
def dilate_labels(
28+
fn: Callable[[np.ndarray], np.ndarray],
29+
iterations: int = 1,
30+
*,
31+
use_gpu: bool = False,
32+
) -> Callable[[np.ndarray], np.ndarray]:
33+
"""Wrap a segmentation callable to grow its labels after each tile.
34+
35+
Applies a single-pass grey dilation to whatever ``fn`` returns, before
36+
``tile_process``/``stage_tile`` trim the overlap halo and merge across
37+
tile boundaries — so dilated labels still stitch correctly at tile
38+
edges.
39+
40+
Parameters
41+
----------
42+
fn : Callable[[np.ndarray], np.ndarray]
43+
Any segmentation function with the ``tile_process``/``stage_tile``
44+
contract (one tile in, integer label array out).
45+
iterations : int, optional
46+
Pixels to grow each label by (grey-dilation footprint size
47+
``2 * iterations + 1``, single pass). Default 1. Values ``<= 0``
48+
disable dilation — ``fn`` is returned unwrapped.
49+
use_gpu : bool, optional
50+
Dilate via cupyx instead of scipy. Independent of whatever backend
51+
``fn`` itself uses internally.
52+
53+
Returns
54+
-------
55+
Callable[[np.ndarray], np.ndarray]
56+
Picklable function ready for ``tile_process``/``stage_tile``. If
57+
``iterations <= 0``, this is ``fn`` itself.
58+
"""
59+
if iterations <= 0:
60+
return fn
61+
return partial(_run, fn=fn, iterations=iterations, use_gpu=use_gpu)
62+
63+
64+
def _run(
65+
block: np.ndarray,
66+
fn: Callable[[np.ndarray], np.ndarray],
67+
iterations: int,
68+
use_gpu: bool,
69+
) -> np.ndarray:
70+
"""Run ``fn`` on ``block``, then grow the resulting labels.
71+
72+
Parameters
73+
----------
74+
block : np.ndarray
75+
One image tile.
76+
fn : Callable[[np.ndarray], np.ndarray]
77+
Segmentation function to run first.
78+
iterations : int
79+
Pixels to grow each label by.
80+
use_gpu : bool
81+
Dilate via cupyx instead of scipy.
82+
83+
Returns
84+
-------
85+
np.ndarray
86+
Dilated integer label array, same shape as ``fn``'s output.
87+
"""
88+
labels = fn(block)
89+
size = 2 * iterations + 1
90+
91+
if use_gpu:
92+
import cupy as cp
93+
from cupyx.scipy.ndimage import grey_dilation
94+
95+
labels = cp.asnumpy(grey_dilation(cp.asarray(labels), size=size))
96+
else:
97+
from scipy.ndimage import grey_dilation
98+
99+
labels = grey_dilation(labels, size=size)
100+
101+
return labels

tests/test_postprocess.py

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,41 @@
1+
"""Self-contained tests for the dilate_labels post-processing wrapper."""
2+
3+
import pickle
4+
5+
import numpy as np
6+
7+
8+
def _make_blob_labels(shape=(1, 64, 64)):
9+
labels = np.zeros(shape, dtype="int32")
10+
labels[0, 28:36, 28:36] = 1
11+
return labels
12+
13+
14+
def test_dilate_labels_grows_mask():
15+
from patchworks import dilate_labels
16+
17+
fn = lambda tile: _make_blob_labels(tile.shape) # noqa: E731
18+
plain = fn(np.zeros((1, 64, 64)))
19+
dilated = dilate_labels(fn, iterations=2)(np.zeros((1, 64, 64)))
20+
21+
assert (dilated > 0).sum() > (plain > 0).sum()
22+
23+
24+
def test_dilate_labels_zero_iterations_is_noop():
25+
from patchworks import dilate_labels
26+
27+
fn = lambda tile: _make_blob_labels(tile.shape) # noqa: E731
28+
29+
assert dilate_labels(fn, iterations=0) is fn
30+
31+
32+
def test_dilate_labels_picklable():
33+
from patchworks.plugins.dog import dog_label_fn
34+
35+
from patchworks import dilate_labels
36+
37+
fn = dilate_labels(
38+
dog_label_fn(low_sigma=1.0, high_sigma=4.0, threshold=0.01),
39+
iterations=2,
40+
)
41+
pickle.loads(pickle.dumps(fn))

0 commit comments

Comments
 (0)