numba-morph 0.1.0__py3-none-any.whl
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.
- numba_morph/__init__.py +15 -0
- numba_morph/_chamfer.py +409 -0
- numba_morph/_edt.py +156 -0
- numba_morph/_scan.py +696 -0
- numba_morph/_watershed.py +165 -0
- numba_morph/cdt.py +82 -0
- numba_morph/dilation.py +42 -0
- numba_morph/erosion.py +102 -0
- numba_morph/fill_holes.py +34 -0
- numba_morph/gradient.py +50 -0
- numba_morph/h_min_max.py +156 -0
- numba_morph/laplace.py +53 -0
- numba_morph/local_min_max.py +83 -0
- numba_morph/open_close.py +88 -0
- numba_morph/reconstruction.py +115 -0
- numba_morph/reconstruction_old.py +423 -0
- numba_morph/top_hats.py +89 -0
- numba_morph/utils.py +111 -0
- numba_morph/watershed.py +68 -0
- numba_morph/welford.py +47 -0
- numba_morph-0.1.0.dist-info/METADATA +132 -0
- numba_morph-0.1.0.dist-info/RECORD +25 -0
- numba_morph-0.1.0.dist-info/WHEEL +5 -0
- numba_morph-0.1.0.dist-info/licenses/LICENSE.md +9 -0
- numba_morph-0.1.0.dist-info/top_level.txt +1 -0
numba_morph/__init__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
__version__ = "0.1.0"
|
|
2
|
+
|
|
3
|
+
from .cdt import distance_transform_cdt
|
|
4
|
+
from .dilation import dilation
|
|
5
|
+
from .erosion import erosion
|
|
6
|
+
from .gradient import morphological_gradient
|
|
7
|
+
from .h_min_max import h_minima, h_maxima
|
|
8
|
+
from .laplace import morphological_laplace
|
|
9
|
+
from .local_min_max import local_minima, local_maxima
|
|
10
|
+
from .open_close import opening, closing
|
|
11
|
+
from .reconstruction import reconstruction
|
|
12
|
+
from .top_hats import white_tophat, black_tophat
|
|
13
|
+
from .utils import generate_sphere_structure
|
|
14
|
+
from .watershed import watershed
|
|
15
|
+
from .welford import welford_mean_std_w_mask
|
numba_morph/_chamfer.py
ADDED
|
@@ -0,0 +1,409 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from numba import njit, prange
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def _generate_offsets(weights):
|
|
6
|
+
"""
|
|
7
|
+
Generate causal and anti-causal offsets for Chamfer masks.
|
|
8
|
+
weights : tuple of ints
|
|
9
|
+
For 2D: (face_weight, edge_weight)
|
|
10
|
+
For 3D: (face_weight, edge_weight, corner_weight)
|
|
11
|
+
Returns two arrays of shape (N, ndim+1) with columns [offset_0, ..., offset_{ndim-1}, weight].
|
|
12
|
+
"""
|
|
13
|
+
ndim = len(weights)
|
|
14
|
+
causal = []
|
|
15
|
+
anti_causal = []
|
|
16
|
+
|
|
17
|
+
if ndim == 2:
|
|
18
|
+
# 2D offsets: (dy, dx)
|
|
19
|
+
for dy in (-1, 0, 1):
|
|
20
|
+
for dx in (-1, 0, 1):
|
|
21
|
+
if dy == 0 and dx == 0:
|
|
22
|
+
continue
|
|
23
|
+
manhattan = abs(dy) + abs(dx)
|
|
24
|
+
if manhattan > ndim: # max Manhattan distance is 2
|
|
25
|
+
continue
|
|
26
|
+
w = weights[manhattan - 1]
|
|
27
|
+
# causal: first non‑zero component is negative
|
|
28
|
+
if dy < 0 or (dy == 0 and dx < 0):
|
|
29
|
+
causal.append((dy, dx, w))
|
|
30
|
+
if dy > 0 or (dy == 0 and dx > 0):
|
|
31
|
+
anti_causal.append((dy, dx, w))
|
|
32
|
+
|
|
33
|
+
elif ndim == 3:
|
|
34
|
+
# 3D offsets: (dz, dy, dx)
|
|
35
|
+
for dz in (-1, 0, 1):
|
|
36
|
+
for dy in (-1, 0, 1):
|
|
37
|
+
for dx in (-1, 0, 1):
|
|
38
|
+
if dz == 0 and dy == 0 and dx == 0:
|
|
39
|
+
continue
|
|
40
|
+
manhattan = abs(dz) + abs(dy) + abs(dx)
|
|
41
|
+
if manhattan > ndim: # max Manhattan distance is 3
|
|
42
|
+
continue
|
|
43
|
+
w = weights[manhattan - 1]
|
|
44
|
+
if dz < 0 or (dz == 0 and dy < 0) or (dz == 0 and dy == 0 and dx < 0):
|
|
45
|
+
causal.append((dz, dy, dx, w))
|
|
46
|
+
if dz > 0 or (dz == 0 and dy > 0) or (dz == 0 and dy == 0 and dx > 0):
|
|
47
|
+
anti_causal.append((dz, dy, dx, w))
|
|
48
|
+
else:
|
|
49
|
+
raise ValueError("Only 2D or 3D are supported")
|
|
50
|
+
|
|
51
|
+
return np.array(causal, dtype=np.int32), np.array(anti_causal, dtype=np.int32)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@njit(cache=True)
|
|
55
|
+
def _apply_pass_2d(input, offsets,
|
|
56
|
+
y_start, y_end, y_step,
|
|
57
|
+
x_start, x_end, x_step,
|
|
58
|
+
H, W, max_val):
|
|
59
|
+
changed = False
|
|
60
|
+
for y in range(y_start, y_end, y_step):
|
|
61
|
+
for x in range(x_start, x_end, x_step):
|
|
62
|
+
if input[y, x] == 0:
|
|
63
|
+
continue
|
|
64
|
+
cur = input[y, x]
|
|
65
|
+
new_val = cur
|
|
66
|
+
for dy, dx, w in offsets:
|
|
67
|
+
ny = y + dy
|
|
68
|
+
nx = x + dx
|
|
69
|
+
if 0 <= ny < H and 0 <= nx < W:
|
|
70
|
+
neigh = input[ny, nx]
|
|
71
|
+
if neigh != max_val:
|
|
72
|
+
cand = neigh + w
|
|
73
|
+
if cand < new_val:
|
|
74
|
+
new_val = cand
|
|
75
|
+
if new_val < cur:
|
|
76
|
+
input[y, x] = new_val
|
|
77
|
+
changed = True
|
|
78
|
+
return changed
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
@njit(parallel=True, cache=True)
|
|
82
|
+
def _apply_pass_2d_batch(input, offsets,
|
|
83
|
+
y_start, y_end, y_step,
|
|
84
|
+
x_start, x_end, x_step,
|
|
85
|
+
L, H, W, max_val):
|
|
86
|
+
changed = 0
|
|
87
|
+
for l in prange(L):
|
|
88
|
+
for y in range(y_start, y_end, y_step):
|
|
89
|
+
for x in range(x_start, x_end, x_step):
|
|
90
|
+
if input[l, y, x] == 0:
|
|
91
|
+
continue
|
|
92
|
+
cur = input[l, y, x]
|
|
93
|
+
new_val = cur
|
|
94
|
+
for dy, dx, w in offsets:
|
|
95
|
+
ny = y + dy
|
|
96
|
+
nx = x + dx
|
|
97
|
+
if 0 <= ny < H and 0 <= nx < W:
|
|
98
|
+
neigh = input[l, ny, nx]
|
|
99
|
+
if neigh != max_val:
|
|
100
|
+
cand = neigh + w
|
|
101
|
+
if cand < new_val:
|
|
102
|
+
new_val = cand
|
|
103
|
+
if new_val < cur:
|
|
104
|
+
input[l, y, x] = new_val
|
|
105
|
+
changed += 1
|
|
106
|
+
return True if changed > 0 else False
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@njit(cache=True)
|
|
110
|
+
def _apply_pass_3d(input, offsets,
|
|
111
|
+
z_start, z_end, z_step,
|
|
112
|
+
y_start, y_end, y_step,
|
|
113
|
+
x_start, x_end, x_step,
|
|
114
|
+
D, H, W, max_val):
|
|
115
|
+
changed = False
|
|
116
|
+
for z in range(z_start, z_end, z_step):
|
|
117
|
+
for y in range(y_start, y_end, y_step):
|
|
118
|
+
for x in range(x_start, x_end, x_step):
|
|
119
|
+
if input[z, y, x] == 0:
|
|
120
|
+
continue
|
|
121
|
+
cur = input[z, y, x]
|
|
122
|
+
new_val = cur
|
|
123
|
+
for dz, dy, dx, w in offsets:
|
|
124
|
+
nz = z + dz
|
|
125
|
+
ny = y + dy
|
|
126
|
+
nx = x + dx
|
|
127
|
+
if 0 <= nz < D and 0 <= ny < H and 0 <= nx < W:
|
|
128
|
+
neigh = input[nz, ny, nx]
|
|
129
|
+
if neigh != max_val:
|
|
130
|
+
cand = neigh + w
|
|
131
|
+
if cand < new_val:
|
|
132
|
+
new_val = cand
|
|
133
|
+
if new_val < cur:
|
|
134
|
+
input[z, y, x] = new_val
|
|
135
|
+
changed = True
|
|
136
|
+
return changed
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
@njit(parallel=True, cache=True)
|
|
140
|
+
def _apply_pass_3d_batch(input, offsets,
|
|
141
|
+
z_start, z_end, z_step,
|
|
142
|
+
y_start, y_end, y_step,
|
|
143
|
+
x_start, x_end, x_step,
|
|
144
|
+
L, D, H, W, max_val):
|
|
145
|
+
changed = 0
|
|
146
|
+
for l in prange(L):
|
|
147
|
+
for z in range(z_start, z_end, z_step):
|
|
148
|
+
for y in range(y_start, y_end, y_step):
|
|
149
|
+
for x in range(x_start, x_end, x_step):
|
|
150
|
+
if input[l, z, y, x] == 0:
|
|
151
|
+
continue
|
|
152
|
+
cur = input[l, z, y, x]
|
|
153
|
+
new_val = cur
|
|
154
|
+
for dz, dy, dx, w in offsets:
|
|
155
|
+
nz = z + dz
|
|
156
|
+
ny = y + dy
|
|
157
|
+
nx = x + dx
|
|
158
|
+
if 0 <= nz < D and 0 <= ny < H and 0 <= nx < W:
|
|
159
|
+
neigh = input[l, nz, ny, nx]
|
|
160
|
+
if neigh != max_val:
|
|
161
|
+
cand = neigh + w
|
|
162
|
+
if cand < new_val:
|
|
163
|
+
new_val = cand
|
|
164
|
+
if new_val < cur:
|
|
165
|
+
input[l, z, y, x] = new_val
|
|
166
|
+
changed += 1
|
|
167
|
+
return True if changed > 0 else False
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
@njit(parallel=True, fastmath=True, cache=True)
|
|
171
|
+
def _chamfer_2d_chunk(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim):
|
|
172
|
+
H, W = input.shape
|
|
173
|
+
chunk_size = (size_of_largest_dim // num_bands) + 1
|
|
174
|
+
|
|
175
|
+
for it in range(size_of_largest_dim):
|
|
176
|
+
changed = np.zeros(num_bands, dtype=np.bool_)
|
|
177
|
+
|
|
178
|
+
for band in prange(num_bands):
|
|
179
|
+
start = band * chunk_size
|
|
180
|
+
end = min(start + chunk_size, size_of_largest_dim)
|
|
181
|
+
|
|
182
|
+
# forward pass
|
|
183
|
+
if _apply_pass_2d(input, causal,
|
|
184
|
+
start, end, 1, # y: start -> end-1
|
|
185
|
+
0, W, 1, # x: 0 -> W-1
|
|
186
|
+
H, W, max_val):
|
|
187
|
+
changed[band] = True
|
|
188
|
+
|
|
189
|
+
# backward pass
|
|
190
|
+
if _apply_pass_2d(input, anti_causal,
|
|
191
|
+
end - 1, start - 1, -1, # y: end-1 down to start
|
|
192
|
+
W - 1, -1, -1, # x: W-1 down to 0
|
|
193
|
+
H, W, max_val):
|
|
194
|
+
changed[band] = True
|
|
195
|
+
|
|
196
|
+
if not np.any(changed):
|
|
197
|
+
break
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
@njit(parallel=True, fastmath=True, cache=True)
|
|
201
|
+
def _chamfer_2d_chunk_batch(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim):
|
|
202
|
+
L, H, W = input.shape
|
|
203
|
+
chunk_size = (size_of_largest_dim // num_bands) + 1
|
|
204
|
+
|
|
205
|
+
for it in range(size_of_largest_dim):
|
|
206
|
+
changed = np.zeros(num_bands, dtype=np.bool_)
|
|
207
|
+
|
|
208
|
+
for band in prange(num_bands):
|
|
209
|
+
start = band * chunk_size
|
|
210
|
+
end = min(start + chunk_size, size_of_largest_dim)
|
|
211
|
+
|
|
212
|
+
# forward pass
|
|
213
|
+
if _apply_pass_2d_batch(input, causal,
|
|
214
|
+
start, end, 1, # y: start -> end-1
|
|
215
|
+
0, W, 1, # x: 0 -> W-1
|
|
216
|
+
L, H, W, max_val):
|
|
217
|
+
changed[band] = True
|
|
218
|
+
|
|
219
|
+
# backward pass
|
|
220
|
+
if _apply_pass_2d_batch(input, anti_causal,
|
|
221
|
+
end - 1, start - 1, -1, # y: end-1 down to start
|
|
222
|
+
W - 1, -1, -1, # x: W-1 down to 0
|
|
223
|
+
L, H, W, max_val):
|
|
224
|
+
changed[band] = True
|
|
225
|
+
|
|
226
|
+
if not np.any(changed):
|
|
227
|
+
break
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
@njit(parallel=True, fastmath=True, cache=True)
|
|
231
|
+
def _chamfer_3d_chunk(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim):
|
|
232
|
+
|
|
233
|
+
D, H, W = input.shape
|
|
234
|
+
chunk_size = (size_of_largest_dim // num_bands) + 1
|
|
235
|
+
|
|
236
|
+
for it in range(size_of_largest_dim):
|
|
237
|
+
changed = np.zeros(num_bands, dtype=np.bool_)
|
|
238
|
+
|
|
239
|
+
for band in prange(num_bands):
|
|
240
|
+
start = band * chunk_size
|
|
241
|
+
end = min(start + chunk_size, size_of_largest_dim)
|
|
242
|
+
|
|
243
|
+
# forward pass
|
|
244
|
+
if _apply_pass_3d(input, causal,
|
|
245
|
+
start, end, 1, # y: start -> end-1
|
|
246
|
+
0, H, 1,
|
|
247
|
+
0, W, 1, # x: 0 -> W-1
|
|
248
|
+
D, H, W, max_val):
|
|
249
|
+
changed[band] = True
|
|
250
|
+
|
|
251
|
+
# backward pass
|
|
252
|
+
if _apply_pass_3d(input, anti_causal,
|
|
253
|
+
end - 1, start - 1, -1, # y: end-1 down to start
|
|
254
|
+
H - 1, -1, -1,
|
|
255
|
+
W - 1, -1, -1, # x: W-1 down to 0
|
|
256
|
+
D, H, W, max_val):
|
|
257
|
+
changed[band] = True
|
|
258
|
+
|
|
259
|
+
if not np.any(changed):
|
|
260
|
+
break
|
|
261
|
+
|
|
262
|
+
|
|
263
|
+
@njit(parallel=True, fastmath=True, cache=True)
|
|
264
|
+
def _chamfer_3d_chunk_batch(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim):
|
|
265
|
+
L, D, H, W = input.shape
|
|
266
|
+
chunk_size = (size_of_largest_dim // num_bands) + 1
|
|
267
|
+
|
|
268
|
+
for it in range(size_of_largest_dim):
|
|
269
|
+
changed = np.zeros(num_bands, dtype=np.bool_)
|
|
270
|
+
|
|
271
|
+
for band in prange(num_bands):
|
|
272
|
+
start = band * chunk_size
|
|
273
|
+
end = min(start + chunk_size, size_of_largest_dim)
|
|
274
|
+
|
|
275
|
+
# forward pass
|
|
276
|
+
if _apply_pass_3d_batch(input, causal,
|
|
277
|
+
start, end, 1,
|
|
278
|
+
0, H, 1,
|
|
279
|
+
0, W, 1, # x: 0 -> W-1
|
|
280
|
+
L, D, H, W, max_val):
|
|
281
|
+
changed[band] = True
|
|
282
|
+
|
|
283
|
+
# backward pass
|
|
284
|
+
if _apply_pass_3d_batch(input, anti_causal,
|
|
285
|
+
end - 1, start - 1, -1,
|
|
286
|
+
H - 1, -1, -1,
|
|
287
|
+
W - 1, -1, -1, # x: W-1 down to 0
|
|
288
|
+
L, D, H, W, max_val):
|
|
289
|
+
changed[band] = True
|
|
290
|
+
|
|
291
|
+
if not np.any(changed):
|
|
292
|
+
break
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def reorder_offset_list(offset_list, spatial_perm):
|
|
296
|
+
new_list = []
|
|
297
|
+
for item in offset_list:
|
|
298
|
+
spatial_offsets = item[:-1] # all but the last element (weight)
|
|
299
|
+
weight = item[-1]
|
|
300
|
+
# Reorder spatial offsets according to spatial_perm
|
|
301
|
+
new_spatial = tuple(spatial_offsets[i] for i in spatial_perm)
|
|
302
|
+
new_list.append((*new_spatial, weight))
|
|
303
|
+
return new_list
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
def _make_ranges(shape):
|
|
307
|
+
"""Return (fwd_ranges, bwd_ranges) for a given spatial shape."""
|
|
308
|
+
fwd = tuple((0, d, 1) for d in shape)
|
|
309
|
+
bwd = tuple((d - 1, -1, -1) for d in shape)
|
|
310
|
+
return fwd, bwd
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
# ---------- Dispatcher for a single forward/backward pass ----------
|
|
314
|
+
def _run_pass(input, causal, anti_causal, ranges_fwd, ranges_bwd, ndim, batch, max_val):
|
|
315
|
+
"""
|
|
316
|
+
Call the appropriate _apply_pass_* function with the given ranges.
|
|
317
|
+
'ranges_fwd' and 'ranges_bwd' are tuples of (start, end, step) per axis.
|
|
318
|
+
"""
|
|
319
|
+
if ndim == 2:
|
|
320
|
+
if batch:
|
|
321
|
+
func = _apply_pass_2d_batch
|
|
322
|
+
# ranges: (y_range, x_range)
|
|
323
|
+
(y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_fwd
|
|
324
|
+
func(input, causal, y_s, y_e, y_st, x_s, x_e, x_st, input.shape[0], input.shape[1], input.shape[2], max_val)
|
|
325
|
+
(y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_bwd
|
|
326
|
+
func(input, anti_causal, y_s, y_e, y_st, x_s, x_e, x_st, input.shape[0], input.shape[1], input.shape[2], max_val)
|
|
327
|
+
else:
|
|
328
|
+
func = _apply_pass_2d
|
|
329
|
+
(y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_fwd
|
|
330
|
+
func(input, causal, y_s, y_e, y_st, x_s, x_e, x_st, input.shape[0], input.shape[1], max_val)
|
|
331
|
+
(y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_bwd
|
|
332
|
+
func(input, anti_causal, y_s, y_e, y_st, x_s, x_e, x_st, input.shape[0], input.shape[1], max_val)
|
|
333
|
+
else: # ndim == 3
|
|
334
|
+
if batch:
|
|
335
|
+
func = _apply_pass_3d_batch
|
|
336
|
+
(z_s, z_e, z_st), (y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_fwd
|
|
337
|
+
func(input, causal, z_s, z_e, z_st, y_s, y_e, y_st, x_s, x_e, x_st,
|
|
338
|
+
input.shape[0], input.shape[1], input.shape[2], input.shape[3], max_val)
|
|
339
|
+
(z_s, z_e, z_st), (y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_bwd
|
|
340
|
+
func(input, anti_causal, z_s, z_e, z_st, y_s, y_e, y_st, x_s, x_e, x_st,
|
|
341
|
+
input.shape[0], input.shape[1], input.shape[2], input.shape[3], max_val)
|
|
342
|
+
else:
|
|
343
|
+
func = _apply_pass_3d
|
|
344
|
+
(z_s, z_e, z_st), (y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_fwd
|
|
345
|
+
func(input, causal, z_s, z_e, z_st, y_s, y_e, y_st, x_s, x_e, x_st,
|
|
346
|
+
input.shape[0], input.shape[1], input.shape[2], max_val)
|
|
347
|
+
(z_s, z_e, z_st), (y_s, y_e, y_st), (x_s, x_e, x_st) = ranges_bwd
|
|
348
|
+
func(input, anti_causal, z_s, z_e, z_st, y_s, y_e, y_st, x_s, x_e, x_st,
|
|
349
|
+
input.shape[0], input.shape[1], input.shape[2], max_val)
|
|
350
|
+
|
|
351
|
+
|
|
352
|
+
# ---------- Dispatcher for the chunked iterative version ----------
|
|
353
|
+
def _run_chunked(input, max_val, num_bands, causal, anti_causal, ndim, batch, size_of_largest_dim):
|
|
354
|
+
"""Call the appropriate _chamfer_*_chunk* function."""
|
|
355
|
+
if ndim == 2:
|
|
356
|
+
if batch:
|
|
357
|
+
_chamfer_2d_chunk_batch(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim)
|
|
358
|
+
else:
|
|
359
|
+
_chamfer_2d_chunk(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim)
|
|
360
|
+
else: # ndim == 3
|
|
361
|
+
if batch:
|
|
362
|
+
_chamfer_3d_chunk_batch(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim)
|
|
363
|
+
else:
|
|
364
|
+
_chamfer_3d_chunk(input, max_val, num_bands, causal, anti_causal, size_of_largest_dim)
|
|
365
|
+
|
|
366
|
+
|
|
367
|
+
# ---------- Helper: permute for chunking ----------
|
|
368
|
+
def _permute_for_chunk(input, causal, anti_causal, working_dim, chunk_dim, batch):
|
|
369
|
+
"""
|
|
370
|
+
Transpose the input so that the chunked axis becomes the first spatial axis,
|
|
371
|
+
and reorder the offsets accordingly. Returns (transposed_input, new_causal, new_anti_causal, perm, inv_perm).
|
|
372
|
+
"""
|
|
373
|
+
spatial_index = working_dim + chunk_dim # absolute axis (0-based among spatial)
|
|
374
|
+
if batch:
|
|
375
|
+
total_axes = working_dim + 1
|
|
376
|
+
abs_chunk_axis = spatial_index + 1 # batch is axis 0
|
|
377
|
+
perm = [0, abs_chunk_axis] + [i for i in range(1, total_axes) if i != abs_chunk_axis]
|
|
378
|
+
spatial_perm = [i - 1 for i in perm[1:]] # spatial order after permutation
|
|
379
|
+
else:
|
|
380
|
+
perm = [spatial_index] + [i for i in range(working_dim) if i != spatial_index]
|
|
381
|
+
spatial_perm = perm
|
|
382
|
+
|
|
383
|
+
new_causal = reorder_offset_list(causal, spatial_perm)
|
|
384
|
+
new_anti_causal = reorder_offset_list(anti_causal, spatial_perm)
|
|
385
|
+
return np.transpose(input, perm), new_causal, new_anti_causal, perm, np.argsort(perm)
|
|
386
|
+
|
|
387
|
+
|
|
388
|
+
# ---------- Refactored _chamfer ----------
|
|
389
|
+
def _chamfer(input, max_val, num_bands, causal, anti_causal, working_dim,
|
|
390
|
+
chunk, size_of_largest_dim, chunk_dim, batch):
|
|
391
|
+
"""
|
|
392
|
+
Compute Chamfer distance transform.
|
|
393
|
+
"""
|
|
394
|
+
if chunk:
|
|
395
|
+
# Move the chunked axis to the front (after batch if present)
|
|
396
|
+
input, causal, anti_causal, perm, inv_perm = _permute_for_chunk(
|
|
397
|
+
input, causal, anti_causal, working_dim, chunk_dim, batch
|
|
398
|
+
)
|
|
399
|
+
# Run the chunked algorithm (the chunk function assumes the first spatial axis is the chunked one)
|
|
400
|
+
_run_chunked(input, max_val, num_bands, causal, anti_causal, working_dim, batch, size_of_largest_dim)
|
|
401
|
+
# Restore original axis order
|
|
402
|
+
input = np.transpose(input, inv_perm)
|
|
403
|
+
else:
|
|
404
|
+
# No chunking: simply run one forward and one backward pass over all axes
|
|
405
|
+
shape = input.shape[1:] if batch else input.shape # spatial shape
|
|
406
|
+
fwd_ranges, bwd_ranges = _make_ranges(shape)
|
|
407
|
+
_run_pass(input, causal, anti_causal, fwd_ranges, bwd_ranges, working_dim, batch, max_val)
|
|
408
|
+
|
|
409
|
+
return input
|
numba_morph/_edt.py
ADDED
|
@@ -0,0 +1,156 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from numba import njit
|
|
3
|
+
|
|
4
|
+
@njit
|
|
5
|
+
def _voronoi_1d(site_pos, site_R, L):
|
|
6
|
+
# Collect valid sites (site_pos != -1)
|
|
7
|
+
# We'll store their indices in arrays f and g (stack)
|
|
8
|
+
f = np.empty(L, dtype=np.int32) # site index (position in array)
|
|
9
|
+
n = 0
|
|
10
|
+
for i in range(L):
|
|
11
|
+
if site_pos[i] != -1:
|
|
12
|
+
f[n] = i
|
|
13
|
+
n += 1
|
|
14
|
+
|
|
15
|
+
best_idx = np.full(L, -1, dtype=np.int32)
|
|
16
|
+
if n == 0:
|
|
17
|
+
return best_idx
|
|
18
|
+
|
|
19
|
+
# Build the lower envelope using a stack
|
|
20
|
+
stack = np.empty(n, dtype=np.int32)
|
|
21
|
+
top = -1
|
|
22
|
+
for i in range(n):
|
|
23
|
+
# While stack has at least 2 sites, check if the last site is obsolete
|
|
24
|
+
while top >= 1:
|
|
25
|
+
# indices of the three sites: prev, last, new
|
|
26
|
+
idx1 = f[stack[top - 1]]
|
|
27
|
+
idx2 = f[stack[top]]
|
|
28
|
+
idx3 = f[i]
|
|
29
|
+
|
|
30
|
+
s1 = site_pos[idx1]
|
|
31
|
+
s2 = site_pos[idx2]
|
|
32
|
+
s3 = site_pos[idx3]
|
|
33
|
+
R1 = site_R[idx1]
|
|
34
|
+
R2 = site_R[idx2]
|
|
35
|
+
R3 = site_R[idx3]
|
|
36
|
+
|
|
37
|
+
# Intersection x between site2 and site3: (x - s2)^2 + R2 = (x - s3)^2 + R3
|
|
38
|
+
# => x23 = (R3 - R2 + s3^2 - s2^2) / (2*(s3 - s2))
|
|
39
|
+
x23 = (R3 - R2 + s3 * s3 - s2 * s2) / (2.0 * (s3 - s2))
|
|
40
|
+
# Intersection between site1 and site2
|
|
41
|
+
x12 = (R2 - R1 + s2 * s2 - s1 * s1) / (2.0 * (s2 - s1))
|
|
42
|
+
|
|
43
|
+
if x23 <= x12:
|
|
44
|
+
# last site (idx2) is never the best, pop it
|
|
45
|
+
top -= 1
|
|
46
|
+
else:
|
|
47
|
+
break
|
|
48
|
+
# Push new site
|
|
49
|
+
top += 1
|
|
50
|
+
stack[top] = i
|
|
51
|
+
|
|
52
|
+
# Now the stack contains the sites forming the lower envelope.
|
|
53
|
+
# Traverse the line and assign the best site.
|
|
54
|
+
l = 0 # pointer into stack
|
|
55
|
+
for x in range(L):
|
|
56
|
+
# Move l forward while the next site gives a smaller distance
|
|
57
|
+
while l < top:
|
|
58
|
+
idx_cur = f[stack[l]]
|
|
59
|
+
idx_next = f[stack[l + 1]]
|
|
60
|
+
s_cur = site_pos[idx_cur]
|
|
61
|
+
s_next = site_pos[idx_next]
|
|
62
|
+
R_cur = site_R[idx_cur]
|
|
63
|
+
R_next = site_R[idx_next]
|
|
64
|
+
|
|
65
|
+
# distance to current site: (x - s_cur)^2 + R_cur
|
|
66
|
+
# distance to next site: (x - s_next)^2 + R_next
|
|
67
|
+
# If next is better (or equal), advance l
|
|
68
|
+
if (x - s_next) ** 2 + R_next <= (x - s_cur) ** 2 + R_cur:
|
|
69
|
+
l += 1
|
|
70
|
+
else:
|
|
71
|
+
break
|
|
72
|
+
best_idx[x] = f[stack[l]]
|
|
73
|
+
|
|
74
|
+
return best_idx
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
@njit
|
|
78
|
+
def edt_feature_transform_2d(binary):
|
|
79
|
+
"""
|
|
80
|
+
Compute the exact Euclidean Feature Transform for a 2D binary image.
|
|
81
|
+
|
|
82
|
+
Parameters
|
|
83
|
+
----------
|
|
84
|
+
binary : 2D int array, shape (H, W)
|
|
85
|
+
1 for foreground, 0 for background.
|
|
86
|
+
|
|
87
|
+
Returns
|
|
88
|
+
-------
|
|
89
|
+
feat_r : 2D int array, shape (H, W)
|
|
90
|
+
Row coordinate of the nearest background pixel for each pixel.
|
|
91
|
+
feat_c : 2D int array, shape (H, W)
|
|
92
|
+
Column coordinate of the nearest background pixel for each pixel.
|
|
93
|
+
For background pixels, feat_r = row, feat_c = col (self).
|
|
94
|
+
"""
|
|
95
|
+
H, W = binary.shape
|
|
96
|
+
feat_r = np.full((H, W), -1, dtype=np.int32)
|
|
97
|
+
feat_c = np.full((H, W), -1, dtype=np.int32)
|
|
98
|
+
|
|
99
|
+
# Initialise: background pixels point to themselves, foreground to -1
|
|
100
|
+
for i in range(H):
|
|
101
|
+
for j in range(W):
|
|
102
|
+
if binary[i, j] == 0:
|
|
103
|
+
feat_r[i, j] = i
|
|
104
|
+
feat_c[i, j] = j
|
|
105
|
+
|
|
106
|
+
# ----- First pass: along rows (axis 0) -----
|
|
107
|
+
# For each column, run 1D Voronoi on the rows.
|
|
108
|
+
# At this stage, accumulated R = 0.
|
|
109
|
+
for j in range(W):
|
|
110
|
+
site_pos = feat_r[:, j] # row coordinate of site (or -1)
|
|
111
|
+
site_R = np.zeros(H, dtype=np.float64) # no previous dims
|
|
112
|
+
best_idx = _voronoi_1d(site_pos, site_R, H)
|
|
113
|
+
for i in range(H):
|
|
114
|
+
if best_idx[i] != -1:
|
|
115
|
+
# best_idx[i] is the row index of the best site in this column
|
|
116
|
+
# The site's coordinates are (feat_r[best_idx[i], j], j)
|
|
117
|
+
feat_r[i, j] = feat_r[best_idx[i], j]
|
|
118
|
+
feat_c[i, j] = j # column is unchanged in this pass
|
|
119
|
+
|
|
120
|
+
# ----- Second pass: along columns (axis 1) -----
|
|
121
|
+
# For each row, run 1D Voronoi on the columns.
|
|
122
|
+
# Now the accumulated R for each site is (row - feat_r[row, col])^2.
|
|
123
|
+
for i in range(H):
|
|
124
|
+
site_pos = feat_c[i, :] # column coordinate of site (or -1)
|
|
125
|
+
# compute R = (i - feat_r[i, j])^2 for each site
|
|
126
|
+
site_R = np.zeros(W, dtype=np.float64)
|
|
127
|
+
for j in range(W):
|
|
128
|
+
if feat_r[i, j] != -1:
|
|
129
|
+
dr = i - feat_r[i, j]
|
|
130
|
+
site_R[j] = dr * dr
|
|
131
|
+
best_idx = _voronoi_1d(site_pos, site_R, W)
|
|
132
|
+
for j in range(W):
|
|
133
|
+
if best_idx[j] != -1:
|
|
134
|
+
# best_idx[j] is the column index of the best site in this row
|
|
135
|
+
# Use that site's feature coordinates
|
|
136
|
+
best_col = best_idx[j]
|
|
137
|
+
feat_r[i, j] = feat_r[i, best_col]
|
|
138
|
+
feat_c[i, j] = feat_c[i, best_col]
|
|
139
|
+
|
|
140
|
+
return feat_r, feat_c
|
|
141
|
+
|
|
142
|
+
if __name__ == '__main__':
|
|
143
|
+
# Example binary image (0 = background, 1 = foreground)
|
|
144
|
+
binary = np.array([[0, 1, 0],
|
|
145
|
+
[1, 1, 1],
|
|
146
|
+
[0, 1, 0]], dtype=np.int32)
|
|
147
|
+
|
|
148
|
+
feat_r, feat_c = edt_feature_transform_2d(binary)
|
|
149
|
+
|
|
150
|
+
# Compute Euclidean distances (squared) using the feature coordinates
|
|
151
|
+
# For each pixel, distance² = (i - feat_r[i,j])² + (j - feat_c[i,j])²
|
|
152
|
+
dist2 = (np.indices(binary.shape)[0] - feat_r) ** 2 + \
|
|
153
|
+
(np.indices(binary.shape)[1] - feat_c) ** 2
|
|
154
|
+
dist = np.sqrt(dist2) # exact Euclidean distance
|
|
155
|
+
|
|
156
|
+
print(dist)
|