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.
@@ -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
@@ -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)