nbmorph 0.2.0__tar.gz → 0.3.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (26) hide show
  1. {nbmorph-0.2.0/src/nbmorph.egg-info → nbmorph-0.3.0}/PKG-INFO +10 -2
  2. {nbmorph-0.2.0 → nbmorph-0.3.0}/README.md +7 -1
  3. {nbmorph-0.2.0 → nbmorph-0.3.0}/pyproject.toml +3 -3
  4. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/__init__.py +2 -2
  5. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/mode.py +209 -92
  6. nbmorph-0.3.0/src/nbmorph/sorting_networks.py +163 -0
  7. {nbmorph-0.2.0 → nbmorph-0.3.0/src/nbmorph.egg-info}/PKG-INFO +10 -2
  8. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph.egg-info/SOURCES.txt +6 -1
  9. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph.egg-info/requires.txt +1 -0
  10. nbmorph-0.3.0/tests/test_minmax.py +26 -0
  11. nbmorph-0.3.0/tests/test_mode.py +144 -0
  12. nbmorph-0.3.0/tests/test_morphology.py +150 -0
  13. nbmorph-0.3.0/tests/test_sorting_networks.py +28 -0
  14. nbmorph-0.3.0/tests/test_zero_edges.py +24 -0
  15. nbmorph-0.2.0/tests/test_core.py +0 -307
  16. {nbmorph-0.2.0 → nbmorph-0.3.0}/LICENSE +0 -0
  17. {nbmorph-0.2.0 → nbmorph-0.3.0}/setup.cfg +0 -0
  18. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/box_kernel.py +0 -0
  19. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/diamond_kernel.py +0 -0
  20. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/minmax.py +0 -0
  21. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/morphology.py +0 -0
  22. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/ops.py +0 -0
  23. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/utils.py +0 -0
  24. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph/zero_edges.py +0 -0
  25. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph.egg-info/dependency_links.txt +0 -0
  26. {nbmorph-0.2.0 → nbmorph-0.3.0}/src/nbmorph.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: nbmorph
3
- Version: 0.2.0
3
+ Version: 0.3.0
4
4
  Summary: A small package with Numba-accelerated morphological operations.
5
5
  Author-email: Marius Causemann <mariusca@simula.no>
6
6
  License: MIT License
@@ -24,6 +24,7 @@ License: MIT License
24
24
  LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
25
25
  OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
26
26
  SOFTWARE.
27
+ Project-URL: homepage, https://github.com/MariusCausemann/nbmorph
27
28
  Classifier: Programming Language :: Python :: 3
28
29
  Classifier: License :: OSI Approved :: MIT License
29
30
  Classifier: Operating System :: OS Independent
@@ -38,6 +39,7 @@ Provides-Extra: test
38
39
  Requires-Dist: pytest; extra == "test"
39
40
  Requires-Dist: pytest-cov; extra == "test"
40
41
  Requires-Dist: fastmorph; extra == "test"
42
+ Requires-Dist: scipy; extra == "test"
41
43
  Dynamic: license-file
42
44
 
43
45
  # nbmorph
@@ -69,7 +71,13 @@ A small, Numba-accelerated Python package for morphological operations on 3D lab
69
71
 
70
72
  * `smooth_labels_spherical`: Smoothes object boundaries by performing an opening followed by a closing.
71
73
 
72
- ![Effect of Morphological Smoothing](img/smoothing_effect.png)
74
+ * **Mode (majority) filters:** Stencil-based mode filters built on top of compile-time sorting networks for fast, fixed-size neighborhoods. Each comes in two variants:
75
+
76
+ * `onlyzero_mode_box` / `onlyzero_mode_diamond`: Fill *only* background (zero) voxels with the mode of their 3x3x3 box or 6-connected diamond neighborhood, leaving labeled voxels unchanged.
77
+
78
+ * `mode_box` / `mode_diamond`: Replace *every* voxel with the mode of its neighborhood (including the center), useful for denoising labeled volumes.
79
+
80
+ ![Effect of Morphological Smoothing](https://github.com/MariusCausemann/nbmorph/raw/main/img/smoothing_effect.png)
73
81
  *Demonstration of the smoothing effect with varying radii and iterations on a sample image. The smoothing is followed by a dilation operation to fill up the empty space.*
74
82
 
75
83
 
@@ -27,7 +27,13 @@ A small, Numba-accelerated Python package for morphological operations on 3D lab
27
27
 
28
28
  * `smooth_labels_spherical`: Smoothes object boundaries by performing an opening followed by a closing.
29
29
 
30
- ![Effect of Morphological Smoothing](img/smoothing_effect.png)
30
+ * **Mode (majority) filters:** Stencil-based mode filters built on top of compile-time sorting networks for fast, fixed-size neighborhoods. Each comes in two variants:
31
+
32
+ * `onlyzero_mode_box` / `onlyzero_mode_diamond`: Fill *only* background (zero) voxels with the mode of their 3x3x3 box or 6-connected diamond neighborhood, leaving labeled voxels unchanged.
33
+
34
+ * `mode_box` / `mode_diamond`: Replace *every* voxel with the mode of its neighborhood (including the center), useful for denoising labeled volumes.
35
+
36
+ ![Effect of Morphological Smoothing](https://github.com/MariusCausemann/nbmorph/raw/main/img/smoothing_effect.png)
31
37
  *Demonstration of the smoothing effect with varying radii and iterations on a sample image. The smoothing is followed by a dilation operation to fill up the empty space.*
32
38
 
33
39
 
@@ -2,7 +2,7 @@
2
2
 
3
3
  [project]
4
4
  name = "nbmorph"
5
- version = "0.2.0"
5
+ version = "0.3.0"
6
6
  authors = [
7
7
  { name="Marius Causemann", email="mariusca@simula.no" },
8
8
  ]
@@ -10,6 +10,7 @@ description = "A small package with Numba-accelerated morphological operations."
10
10
  readme = "README.md"
11
11
  license = { file="LICENSE" }
12
12
  requires-python = ">=3.8"
13
+ urls = { homepage = "https://github.com/MariusCausemann/nbmorph" }
13
14
  classifiers = [
14
15
  "Programming Language :: Python :: 3",
15
16
  "License :: OSI Approved :: MIT License",
@@ -24,7 +25,7 @@ dependencies = [
24
25
 
25
26
  [project.optional-dependencies]
26
27
  test = [
27
- "pytest", "pytest-cov", "fastmorph",
28
+ "pytest", "pytest-cov", "fastmorph", "scipy",
28
29
  ]
29
30
 
30
31
  [tool.coverage.run]
@@ -32,4 +33,3 @@ branch = false
32
33
 
33
34
  [tool.coverage.report]
34
35
  ignore_errors = true
35
-
@@ -5,10 +5,10 @@ from .morphology import (
5
5
  close_labels_spherical,
6
6
  smooth_labels_spherical,
7
7
  )
8
- from .mode import onlyzero_mode_box,onlyzero_mode_diamond, fast_mode
8
+ from .mode import onlyzero_mode_box, onlyzero_mode_diamond, fast_mode, mode_box, mode_diamond
9
9
  from .minmax import minimum_box, maximum_box, minimum_diamond, maximum_diamond
10
10
  from .zero_edges import zero_label_edges_box, zero_label_edges_diamond
11
11
  from .utils import cycle
12
12
 
13
13
  # Define the package version
14
- __version__ = "0.1.0"
14
+ __version__ = "0.3.0"
@@ -1,13 +1,14 @@
1
1
  import numba
2
2
  import numpy as np
3
3
  from .minmax import maximum_box, maximum_diamond
4
+ from .sorting_networks import sort6_network, sort7_network, sort26_network, sort27_network
4
5
 
5
6
  @numba.njit
6
7
  def fast_mode(a):
7
8
  """
8
9
  Find the mode of a 1D array, ignoring zeros.
9
10
 
10
- This is an O(n^2) algorithm, but fast on small data (len(a) < 50), as needed here.
11
+ This is an O(n^2) algorithm, but fast on small data (len(a) < 50).
11
12
 
12
13
  Args:
13
14
  a (np.ndarray): The input 1D array.
@@ -22,8 +23,7 @@ def fast_modeN(a, N):
22
23
  """
23
24
  Find the mode of the first N elements of a 1D array, ignoring zeros.
24
25
 
25
- This is an O(n^2) algorithm, but fast on small data (N < 50), as needed here.
26
-
26
+ This is an O(n^2) algorithm, but fast on small data (N < 50).
27
27
  Args:
28
28
  a (np.ndarray): The input 1D array.
29
29
  N (int): The number of elements to consider.
@@ -51,99 +51,59 @@ def fast_modeN(a, N):
51
51
  mode = a[i]
52
52
  return mode
53
53
 
54
- @numba.njit(inline="always")
55
- def _cs(a, b):
56
- """
57
- Performs a compare-swap on two values.
58
-
59
- Args:
60
- a: First value.
61
- b: Second value.
62
-
63
- Returns:
64
- Tuple containing the smaller value followed by the larger value.
65
- """
66
- if a > b:
67
- return b, a
68
- else:
69
- return a, b
70
-
71
54
 
72
55
  @numba.njit(inline="always")
73
- def sort26_network(
74
- v0, v1, v2, v3, v4, v5, v6, v7, v8, v9, v10, v11, v12, v13,
75
- v14, v15, v16, v17, v18, v19, v20, v21, v22, v23, v24, v25
76
- ):
56
+ def outer_mode_diamond_kernel(data, z, y, x):
77
57
  """
78
- Sorts 26 elements using a pre-defined sorting network.
58
+ Calculates the mode of a diamond neighborhood in a 3D array.
79
59
 
80
- This function implements a sorting network from https://bertdobbelaere.github.io/sorting_networks.html
81
- to efficiently sort a fixed number of elements.
60
+ The diamond neighborhood includes the 6 direct (6-connected) neighbors of the center point:
61
+ (z, y, x-1), (z, y, x+1), (z, y-1, x), (z, y+1, x), (z-1, y, x), (z+1, y, x).
82
62
 
83
63
  Args:
84
- v0-v25: The 26 values to be sorted.
64
+ data (np.ndarray): The 3D input array.
65
+ z (int): Z-coordinate of the center point.
66
+ y (int): Y-coordinate of the center point.
67
+ x (int): X-coordinate of the center point.
85
68
 
86
69
  Returns:
87
- Tuple containing the 26 input values in sorted order (ascending).
70
+ The mode (most frequent value) of the diamond neighborhood, ignoring zeros.
88
71
  """
89
- v0, v1 = _cs(v0, v1); v2, v3 = _cs(v2, v3); v4, v5 = _cs(v4, v5); v6, v7 = _cs(v6, v7); v8, v9 = _cs(v8, v9); v10, v11 = _cs(v10, v11); v12, v13 = _cs(v12, v13); v14, v15 = _cs(v14, v15); v16, v17 = _cs(v16, v17); v18, v19 = _cs(v18, v19); v20, v21 = _cs(v20, v21); v22, v23 = _cs(v22, v23); v24, v25 = _cs(v24, v25)
90
- v0, v2 = _cs(v0, v2); v1, v3 = _cs(v1, v3); v4, v6 = _cs(v4, v6); v5, v7 = _cs(v5, v7); v8, v10 = _cs(v8, v10); v9, v11 = _cs(v9, v11); v14, v16 = _cs(v14, v16); v15, v17 = _cs(v15, v17); v18, v20 = _cs(v18, v20); v19, v21 = _cs(v19, v21); v22, v24 = _cs(v22, v24); v23, v25 = _cs(v23, v25)
91
- v0, v4 = _cs(v0, v4); v1, v6 = _cs(v1, v6); v2, v5 = _cs(v2, v5); v3, v7 = _cs(v3, v7); v8, v14 = _cs(v8, v14); v9, v16 = _cs(v9, v16); v10, v15 = _cs(v10, v15); v11, v17 = _cs(v11, v17); v18, v22 = _cs(v18, v22); v19, v24 = _cs(v19, v24); v20, v23 = _cs(v20, v23); v21, v25 = _cs(v21, v25)
92
- v0, v18 = _cs(v0, v18); v1, v19 = _cs(v1, v19); v2, v20 = _cs(v2, v20); v3, v21 = _cs(v3, v21); v4, v22 = _cs(v4, v22); v5, v23 = _cs(v5, v23); v6, v24 = _cs(v6, v24); v7, v25 = _cs(v7, v25); v9, v12 = _cs(v9, v12); v13, v16 = _cs(v13, v16)
93
- v3, v11 = _cs(v3, v11); v8, v9 = _cs(v8, v9); v10, v13 = _cs(v10, v13); v12, v15 = _cs(v12, v15); v14, v22 = _cs(v14, v22); v16, v17 = _cs(v16, v17)
94
- v0, v8 = _cs(v0, v8); v1, v9 = _cs(v1, v9); v2, v14 = _cs(v2, v14); v6, v12 = _cs(v6, v12); v7, v15 = _cs(v7, v15); v10, v18 = _cs(v10, v18); v11, v23 = _cs(v11, v23); v13, v19 = _cs(v13, v19); v16, v24 = _cs(v16, v24); v17, v25 = _cs(v17, v25)
95
- v1, v2 = _cs(v1, v2); v3, v18 = _cs(v3, v18); v4, v8 = _cs(v4, v8); v7, v22 = _cs(v7, v22); v17, v21 = _cs(v17, v21); v23, v24 = _cs(v23, v24)
96
- v3, v14 = _cs(v3, v14); v4, v10 = _cs(v4, v10); v5, v18 = _cs(v5, v18); v7, v20 = _cs(v7, v20); v8, v13 = _cs(v8, v13); v11, v22 = _cs(v11, v22); v12, v17 = _cs(v12, v17); v15, v21 = _cs(v15, v21)
97
- v1, v4 = _cs(v1, v4); v5, v6 = _cs(v5, v6); v7, v9 = _cs(v7, v9); v8, v10 = _cs(v8, v10); v15, v17 = _cs(v15, v17); v16, v18 = _cs(v16, v18); v19, v20 = _cs(v19, v20); v21, v24 = _cs(v21, v24)
98
- v2, v5 = _cs(v2, v5); v3, v10 = _cs(v3, v10); v6, v14 = _cs(v6, v14); v9, v13 = _cs(v9, v13); v11, v19 = _cs(v11, v19); v12, v16 = _cs(v12, v16); v15, v22 = _cs(v15, v22); v20, v23 = _cs(v20, v23)
99
- v2, v8 = _cs(v2, v8); v5, v7 = _cs(v5, v7); v6, v9 = _cs(v6, v9); v11, v12 = _cs(v11, v12); v13, v14 = _cs(v13, v14); v16, v19 = _cs(v16, v19); v17, v23 = _cs(v17, v23); v18, v20 = _cs(v18, v20)
100
- v2, v4 = _cs(v2, v4); v3, v5 = _cs(v3, v5); v6, v11 = _cs(v6, v11); v7, v10 = _cs(v7, v10); v9, v16 = _cs(v9, v16); v12, v13 = _cs(v12, v13); v14, v19 = _cs(v14, v19); v15, v18 = _cs(v15, v18); v20, v22 = _cs(v20, v22); v21, v23 = _cs(v21, v23)
101
- v3, v4 = _cs(v3, v4); v5, v8 = _cs(v5, v8); v6, v7 = _cs(v6, v7); v9, v11 = _cs(v9, v11); v10, v12 = _cs(v10, v12); v13, v15 = _cs(v13, v15); v14, v16 = _cs(v14, v16); v17, v20 = _cs(v17, v20); v18, v19 = _cs(v18, v19); v21, v22 = _cs(v21, v22)
102
- v5, v6 = _cs(v5, v6); v7, v8 = _cs(v7, v8); v9, v10 = _cs(v9, v10); v11, v12 = _cs(v11, v12); v13, v14 = _cs(v13, v14); v15, v16 = _cs(v15, v16); v17, v18 = _cs(v17, v18); v19, v20 = _cs(v19, v20)
103
- v4, v5 = _cs(v4, v5); v6, v7 = _cs(v6, v7); v8, v9 = _cs(v8, v9); v10, v11 = _cs(v10, v11); v12, v13 = _cs(v12, v13); v14, v15 = _cs(v14, v15); v16, v17 = _cs(v16, v17); v18, v19 = _cs(v18, v19); v20, v21 = _cs(v20, v21)
104
-
105
- return v0, v1, v2, v3, v4, v5, v6, v7, v8, v9, v10, v11, v12, v13, v14, v15, v16, v17, v18, v19, v20, v21, v22, v23, v24, v25
106
72
 
107
- @numba.njit(inline="always")
108
- def sort6_network(v0, v1, v2, v3, v4, v5):
109
- """
110
- Sorts 6 elements using a pre-defined sorting network.
73
+ (v0, v1, v2, v3, v4, v5) = sort6_network(
74
+ data[z, y, x-1], data[z, y, x+1],
75
+ data[z, y-1, x], data[z, y+1, x],
76
+ data[z-1, y, x], data[z+1, y, x]
77
+ )
111
78
 
112
- This function implements a sorting network from https://bertdobbelaere.github.io/sorting_networks.html
113
- to efficiently sort a fixed number of elements.
79
+ one = np.uint8(1)
80
+ l0 = one
81
+ l1 = (l0 + one) if v1 == v0 and v1 > 0 else one
82
+ l2 = (l1 + one) if v2 == v1 and v2 > 0 else one
83
+ l3 = (l2 + one) if v3 == v2 and v3 > 0 else one
84
+ l4 = (l3 + one) if v4 == v3 and v4 > 0 else one
85
+ l5 = (l4 + one) if v5 == v4 and v5 > 0 else one
114
86
 
115
- Args:
116
- v0-v5: The 6 values to be sorted.
87
+ def _update_max(len1, val1, len2, val2):
88
+ if len2 >= len1:
89
+ return len2, val2
90
+ return len1, val1
117
91
 
118
- Returns:
119
- Tuple containing the 6 input values in sorted order (ascending).
120
- """
121
- v0, v5 = _cs(v0, v5)
122
- v1, v3 = _cs(v1, v3)
123
- v2, v4 = _cs(v2, v4)
124
-
125
- v1, v2 = _cs(v1, v2)
126
- v3, v4 = _cs(v3, v4)
127
-
128
- v0, v3 = _cs(v0, v3)
129
- v2, v5 = _cs(v2, v5)
130
-
131
- v0, v1 = _cs(v0, v1)
132
- v2, v3 = _cs(v2, v3)
133
- v4, v5 = _cs(v4, v5)
134
-
135
- v1, v2 = _cs(v1, v2)
136
- v3, v4 = _cs(v3, v4)
92
+ (l_max, v_mode) = _update_max(l0, v0, l1, v1)
93
+ (l_max, v_mode) = _update_max(l_max, v_mode, l2, v2)
94
+ (l_max, v_mode) = _update_max(l_max, v_mode, l3, v3)
95
+ (l_max, v_mode) = _update_max(l_max, v_mode, l4, v4)
96
+ (l_max, v_mode) = _update_max(l_max, v_mode, l5, v5)
137
97
 
138
- return v0, v1, v2, v3, v4, v5
98
+ return v_mode
99
+
139
100
 
140
101
  @numba.njit(inline="always")
141
- def mode_diamond(data, z, y, x):
102
+ def mode_diamond_kernel(data, z, y, x):
142
103
  """
143
104
  Calculates the mode of a diamond neighborhood in a 3D array.
144
105
 
145
- The diamond neighborhood includes the 6 direct (6-connected) neighbors of the center point:
146
- (z, y, x-1), (z, y, x+1), (z, y-1, x), (z, y+1, x), (z-1, y, x), (z+1, y, x).
106
+ The diamond neighborhood includes the 6 direct (6-connected) neighbors and the center point.
147
107
 
148
108
  Args:
149
109
  data (np.ndarray): The 3D input array.
@@ -155,9 +115,9 @@ def mode_diamond(data, z, y, x):
155
115
  The mode (most frequent value) of the diamond neighborhood, ignoring zeros.
156
116
  """
157
117
 
158
- (v0, v1, v2, v3, v4, v5) = sort6_network(
118
+ (v0, v1, v2, v3, v4, v5, v6) = sort7_network(
159
119
  data[z, y, x-1], data[z, y, x+1],
160
- data[z, y-1, x], data[z, y+1, x],
120
+ data[z, y-1, x], data[z, y, x], data[z, y+1, x],
161
121
  data[z-1, y, x], data[z+1, y, x]
162
122
  )
163
123
 
@@ -168,6 +128,7 @@ def mode_diamond(data, z, y, x):
168
128
  l3 = (l2 + one) if v3 == v2 and v3 > 0 else one
169
129
  l4 = (l3 + one) if v4 == v3 and v4 > 0 else one
170
130
  l5 = (l4 + one) if v5 == v4 and v5 > 0 else one
131
+ l6 = (l5 + one) if v6 == v5 and v6 > 0 else one
171
132
 
172
133
  def _update_max(len1, val1, len2, val2):
173
134
  if len2 >= len1:
@@ -179,13 +140,13 @@ def mode_diamond(data, z, y, x):
179
140
  (l_max, v_mode) = _update_max(l_max, v_mode, l3, v3)
180
141
  (l_max, v_mode) = _update_max(l_max, v_mode, l4, v4)
181
142
  (l_max, v_mode) = _update_max(l_max, v_mode, l5, v5)
143
+ (l_max, v_mode) = _update_max(l_max, v_mode, l6, v6)
182
144
 
183
- #print(v_mode)
184
145
  return v_mode
185
146
 
186
147
 
187
148
  @numba.njit(inline="always")
188
- def mode_box(data, z, y, x):
149
+ def outer_mode_box_kernel(data, z, y, x):
189
150
  """
190
151
  Calculates the mode of a 3x3x3 neighborhood in a 3D array.
191
152
 
@@ -291,14 +252,124 @@ def mode_box(data, z, y, x):
291
252
  l25, v25 = _update_max(l15, v15, l25, v25)
292
253
  return v25
293
254
 
255
+ import numpy as np
256
+ import numba
257
+
258
+ @numba.njit(inline="always")
259
+ def mode_box_kernel(data, z, y, x):
260
+ """
261
+ Calculates the mode of a 3x3x3 neighborhood in a 3D array.
262
+
263
+ The neighborhood includes all 27 surrounding voxels, including the center point.
264
+
265
+ Args:
266
+ data (np.ndarray): The 3D input array.
267
+ z (int): Z-coordinate of the center point.
268
+ y (int): Y-coordinate of the center point.
269
+ x (int): X-coordinate of the center point.
270
+
271
+ Returns:
272
+ The mode (most frequent value) of the 3x3x3 neighborhood, ignoring zeros.
273
+ """
274
+
275
+ (v0, v1, v2, v3, v4, v5, v6, v7, v8, v9, v10, v11,
276
+ v12, v13, v14, v15, v16, v17, v18, v19, v20,
277
+ v21, v22, v23, v24, v25, v26 ) = sort27_network(
278
+ # --- Top Slice (z-1) ---
279
+ data[z-1, y-1, x-1], data[z-1, y-1, x], data[z-1, y-1, x+1],
280
+ data[z-1, y, x-1], data[z-1, y, x], data[z-1, y, x+1],
281
+ data[z-1, y+1, x-1], data[z-1, y+1, x], data[z-1, y+1, x+1],
282
+
283
+ # --- Middle Slice (z) ---
284
+ data[z, y-1, x-1], data[z, y-1, x], data[z, y-1, x+1],
285
+ data[z, y, x-1], data[z, y, x], data[z, y, x+1],
286
+ data[z, y+1, x-1], data[z, y+1, x], data[z, y+1, x+1],
287
+
288
+ # --- Bottom Slice (z+1) ---
289
+ data[z+1, y-1, x-1], data[z+1, y-1, x], data[z+1, y-1, x+1],
290
+ data[z+1, y, x-1], data[z+1, y, x], data[z+1, y, x+1],
291
+ data[z+1, y+1, x-1], data[z+1, y+1, x], data[z+1, y+1, x+1]
292
+ )
293
+
294
+ one = np.uint8(1)
295
+ l0 = one
296
+ l1 = (l0 + one) if v1 == v0 and v1>0 else one
297
+ l2 = (l1 + one) if v2 == v1 and v2>0 else one
298
+ l3 = (l2 + one) if v3 == v2 and v3>0 else one
299
+ l4 = (l3 + one) if v4 == v3 and v4>0 else one
300
+ l5 = (l4 + one) if v5 == v4 and v5>0 else one
301
+ l6 = (l5 + one) if v6 == v5 and v6>0 else one
302
+ l7 = (l6 + one) if v7 == v6 and v7>0 else one
303
+ l8 = (l7 + one) if v8 == v7 and v8>0 else one
304
+ l9 = (l8 + one) if v9 == v8 and v9>0 else one
305
+ l10 = (l9 + one) if v10 == v9 and v10>0 else one
306
+ l11 = (l10 + one) if v11 == v10 and v11>0 else one
307
+ l12 = (l11 + one) if v12 == v11 and v12>0 else one
308
+ l13 = (l12 + one) if v13 == v12 and v13>0 else one
309
+ l14 = (l13 + one) if v14 == v13 and v14>0 else one
310
+ l15 = (l14 + one) if v15 == v14 and v15>0 else one
311
+ l16 = (l15 + one) if v16 == v15 and v16>0 else one
312
+ l17 = (l16 + one) if v17 == v16 and v17>0 else one
313
+ l18 = (l17 + one) if v18 == v17 and v18>0 else one
314
+ l19 = (l18 + one) if v19 == v18 and v19>0 else one
315
+ l20 = (l19 + one) if v20 == v19 and v20>0 else one
316
+ l21 = (l20 + one) if v21 == v20 and v21>0 else one
317
+ l22 = (l21 + one) if v22 == v21 and v22>0 else one
318
+ l23 = (l22 + one) if v23 == v22 and v23>0 else one
319
+ l24 = (l23 + one) if v24 == v23 and v24>0 else one
320
+ l25 = (l24 + one) if v25 == v24 and v25>0 else one
321
+ l26 = (l25 + one) if v26 == v25 and v26>0 else one
322
+
323
+ def _update_max(len1, val1, len2, val2):
324
+ if len2 >= len1:
325
+ return len2, val2
326
+ return len1, val1
327
+
328
+ # Layer 1: 13 parallel comparisons (v26 carries over)
329
+ l1, v1 = _update_max(l0, v0, l1, v1)
330
+ l3, v3 = _update_max(l2, v2, l3, v3)
331
+ l5, v5 = _update_max(l4, v4, l5, v5)
332
+ l7, v7 = _update_max(l6, v6, l7, v7)
333
+ l9, v9 = _update_max(l8, v8, l9, v9)
334
+ l11, v11 = _update_max(l10, v10, l11, v11)
335
+ l13, v13 = _update_max(l12, v12, l13, v13)
336
+ l15, v15 = _update_max(l14, v14, l15, v15)
337
+ l17, v17 = _update_max(l16, v16, l17, v17)
338
+ l19, v19 = _update_max(l18, v18, l19, v19)
339
+ l21, v21 = _update_max(l20, v20, l21, v21)
340
+ l23, v23 = _update_max(l22, v22, l23, v23)
341
+ l25, v25 = _update_max(l24, v24, l25, v25)
342
+
343
+ # Layer 2: Winners from Layer 1 compete (7 parallel comparisons)
344
+ l3, v3 = _update_max(l1, v1, l3, v3)
345
+ l7, v7 = _update_max(l5, v5, l7, v7)
346
+ l11, v11 = _update_max(l9, v9, l11, v11)
347
+ l15, v15 = _update_max(l13, v13, l15, v15)
348
+ l19, v19 = _update_max(l17, v17, l19, v19)
349
+ l23, v23 = _update_max(l21, v21, l23, v23)
350
+ l26, v26 = _update_max(l25, v25, l26, v26)
351
+
352
+ # Layer 3: 3 parallel comparisons (v26 carries over)
353
+ l7, v7 = _update_max(l3, v3, l7, v7)
354
+ l15, v15 = _update_max(l11, v11, l15, v15)
355
+ l23, v23 = _update_max(l19, v19, l23, v23)
356
+
357
+ # Layer 4: 2 parallel comparisons
358
+ l15, v15 = _update_max(l7, v7, l15, v15)
359
+ l26, v26 = _update_max(l23, v23, l26, v26)
360
+
361
+ # Layer 5: Final comparison
362
+ l26, v26 = _update_max(l15, v15, l26, v26)
363
+
364
+ return v26
365
+
294
366
 
295
367
  @numba.njit(inline="always")
296
368
  def load_box_stencil(data, z, y, x, sz, sy, sx, nbs):
297
369
  """
298
370
  Loads a 3x3x3 box stencil into a neighbors array, counting only non-zero values.
299
-
300
- This function extracts the 26 neighbors of a center point in a 3D array, handling
301
- boundary conditions appropriately.
371
+ The stencil includes the center point and its 26 direct neighbors.
372
+ Boundary conditions are handled by checking array dimensions.
302
373
 
303
374
  Args:
304
375
  data (np.ndarray): The 3D input array.
@@ -330,7 +401,7 @@ def load_box_stencil(data, z, y, x, sz, sy, sx, nbs):
330
401
  def load_diamond_stencil(data, z, y, x, sz, sy, sx, nbs):
331
402
  """
332
403
  Loads a diamond stencil into a neighbors array, counting only non-zero values.
333
- The stencil includes the 6 direct neighbors.
404
+ The stencil includes the centerpoint and its 6 direct neighbors.
334
405
  Boundary conditions are handled by checking array dimensions.
335
406
 
336
407
  Parameters:
@@ -343,7 +414,12 @@ def load_diamond_stencil(data, z, y, x, sz, sy, sx, nbs):
343
414
  - The number of non-zero values loaded.
344
415
  """
345
416
  nnz = 0
346
-
417
+
418
+ val = data[z, y, x]
419
+ nbs[nnz] = val
420
+ if val > 0:
421
+ nnz += 1
422
+
347
423
  if z > 0:
348
424
  val = data[z - 1, y, x]
349
425
  nbs[nnz] = val
@@ -382,12 +458,12 @@ def load_diamond_stencil(data, z, y, x, sz, sy, sx, nbs):
382
458
  return nnz
383
459
 
384
460
  @numba.njit
385
- def _mode_borders(data, out, stencil):
461
+ def _mode_borders(data, out, stencil, onlyzero=True):
386
462
  sz, sy, sx = data.shape
387
463
  nbs = np.empty(17, dtype=data.dtype)
388
464
 
389
465
  def process_point(z, y, x):
390
- if data[z, y, x] > 0:
466
+ if onlyzero and data[z, y, x] > 0:
391
467
  out[z, y, x] = data[z, y, x]
392
468
  else:
393
469
  if stencil=="box":
@@ -396,6 +472,7 @@ def _mode_borders(data, out, stencil):
396
472
  nnz = load_diamond_stencil(data, z, y, x, sz,sy,sx, nbs)
397
473
 
398
474
  out[z, y, x] = fast_modeN(nbs, nnz) * (nnz > 0)
475
+
399
476
 
400
477
  # 1. Top and Bottom faces (Z-axis)
401
478
  for z in [0, sz - 1]:
@@ -429,10 +506,30 @@ def _onlyzero_mode_box(data, out=None):
429
506
  if data[z,y,x]>0:
430
507
  out[z,y,x] = data[z,y,x]
431
508
  else:
432
- out[z,y,x] = mode_box(data, z,y,x)
433
- _mode_borders(data, out, stencil="box")
509
+ out[z,y,x] = outer_mode_box_kernel(data, z,y,x)
510
+ _mode_borders(data, out, stencil="box", onlyzero=True)
511
+ return out
512
+
513
+ @numba.njit(parallel=True)
514
+ def _mode_box(data, out=None):
515
+ sz, sy, sx = data.shape
516
+ if out is None:
517
+ out = np.empty_like(data)
518
+ assert data.shape == out.shape
519
+ for z in numba.prange(1, sz-1):
520
+ for y in range(1, sy-1):
521
+ for x in range(1, sx-1):
522
+ out[z,y,x] = mode_box_kernel(data, z,y,x)
523
+ # FIXME: wrong border handling?
524
+ _mode_borders(data, out, stencil="box", onlyzero=False)
434
525
  return out
435
526
 
527
+ @numba.njit(cache=True)
528
+ def mode_box(data, out=None):
529
+ if isinstance(data.dtype.type(0), bool):
530
+ return maximum_box(data, out, onlyzero=False)
531
+ else:
532
+ return _mode_box(data, out=out)
436
533
 
437
534
  @numba.njit(cache=True)
438
535
  def onlyzero_mode_box(data, out=None):
@@ -453,8 +550,8 @@ def _onlyzero_mode_diamond(data, out=None):
453
550
  if data[z,y,x]>0:
454
551
  out[z,y,x] = data[z,y,x]
455
552
  else:
456
- out[z,y,x] = mode_diamond(data, z,y,x)
457
- _mode_borders(data, out, stencil="diamond")
553
+ out[z,y,x] = outer_mode_diamond_kernel(data, z,y,x)
554
+ _mode_borders(data, out, stencil="diamond", onlyzero=True)
458
555
  return out
459
556
 
460
557
  @numba.njit(cache=True)
@@ -463,3 +560,23 @@ def onlyzero_mode_diamond(data, out=None):
463
560
  return maximum_diamond(data, out, onlyzero=True)
464
561
  else:
465
562
  return _onlyzero_mode_diamond(data, out=out)
563
+
564
+ @numba.njit(parallel=True)
565
+ def _mode_diamond(data, out=None):
566
+ sz, sy, sx = data.shape
567
+ if out is None:
568
+ out = np.empty_like(data)
569
+ assert data.shape == out.shape
570
+ for z in numba.prange(1, sz-1):
571
+ for y in range(1, sy-1):
572
+ for x in range(1, sx-1):
573
+ out[z,y,x] = mode_diamond_kernel(data, z,y,x)
574
+ _mode_borders(data, out, stencil="diamond", onlyzero=False)
575
+ return out
576
+
577
+ @numba.njit(cache=True)
578
+ def mode_diamond(data, out=None):
579
+ if isinstance(data.dtype.type(0), bool):
580
+ return maximum_diamond(data, out, onlyzero=False)
581
+ else:
582
+ return _mode_diamond(data, out=out)