From 80d76d25abfd333b98c4a67000f6b7cb26a25c8e Mon Sep 17 00:00:00 2001 From: Jeff Gostick Date: Mon, 14 Sep 2026 00:04:05 -0400 Subject: [PATCH] added merged raster scanning to lt --- src/porespy/filters/_lt_methods.py | 101 +++------- src/porespy/simulations/_tools.py | 251 +---------------------- src/porespy/tools/_sphere_insertions.py | 258 +++++++++++++++++++++++- test/unit/test_simulations_ibop.py | 40 ++++ 4 files changed, 329 insertions(+), 321 deletions(-) diff --git a/src/porespy/filters/_lt_methods.py b/src/porespy/filters/_lt_methods.py index d4659a06e..3b3d26161 100644 --- a/src/porespy/filters/_lt_methods.py +++ b/src/porespy/filters/_lt_methods.py @@ -8,8 +8,9 @@ from skimage.morphology import ball, disk, footprint_rectangle from porespy.tools import ( - _get_axial_extent, _get_uint_dtype, + _insert_disks_at_indices_parallel, + _insert_disks_at_indices_parallel_direct, _make_axial_extent_lookup, get_edt, get_tqdm, @@ -256,73 +257,6 @@ def _remove_lt_contained_disks(indices, previous): return indices[:count] -@njit(parallel=True) -def _insert_lt_disks_at_indices( - lt, - indices, - radius, - value, - ceil_distance, - smooth, -): # pragma: no cover - """Rasterize one equal-radius bucket directly from flat center indices.""" - npts = len(indices) - radius_squared = radius * radius - if lt.ndim == 2: - xlim, ylim = lt.shape - lt_flat = lt.reshape(lt.size) - for q in prange(npts): - ind = indices[q] - i = ind // ylim - j = ind - i * ylim - for x in range(max(0, i - radius), min(i + radius + 1, xlim)): - dx = x - i - y_extent = _get_axial_extent( - radius_squared - dx * dx, - ceil_distance, - smooth, - ) - start = x * ylim + max(0, j - y_extent) - stop = x * ylim + min(j + y_extent + 1, ylim) - for p in range(start, stop): - if lt_flat[p] == 0: - lt_flat[p] = value - elif lt.ndim == 3: - xlim, ylim, zlim = lt.shape - stride0 = ylim * zlim - lt_flat = lt.reshape(lt.size) - for q in prange(npts): - ind = indices[q] - i = ind // stride0 - rem = ind - i * stride0 - j = rem // zlim - k = rem - j * zlim - for x in range(max(0, i - radius), min(i + radius + 1, xlim)): - dx = x - i - dx_squared = dx * dx - yz_extent = _get_axial_extent( - radius_squared - dx_squared, - ceil_distance, - smooth, - ) - for y in range( - max(0, j - yz_extent), - min(j + yz_extent + 1, ylim), - ): - dy = y - j - z_extent = _get_axial_extent( - radius_squared - dx_squared - dy * dy, - ceil_distance, - smooth, - ) - start = (x * ylim + y) * zlim + max(0, k - z_extent) - stop = (x * ylim + y) * zlim + min(k + z_extent + 1, zlim) - for p in range(start, stop): - if lt_flat[p] == 0: - lt_flat[p] = value - return lt - - def local_thickness( im: npt.NDArray, dt: npt.NDArray = None, @@ -508,7 +442,8 @@ def local_thickness_bf( Consecutive radii rasterize only the new distance-transform shell and discard spheres contained by the preceding shell. When requested radii contain gaps, the threshold is filled directly and spheres are rasterized only from its - interface. Both paths draw in parallel from compact flat indices. + interface. Shells write labels directly, while interfaces build a cumulative + boolean union using adaptive direct or merged parallel scan-line writes. Examples -------- @@ -526,6 +461,7 @@ def local_thickness_bf( # inserting at every eligible voxel for every radius. radii = _parse_integer_radii(sizes=sizes, dt=dt, im=im) lt, values = _make_lt_result(im.shape, radii, return_indices) + nwp = np.zeros(im.shape, dtype=bool) interface = np.empty(im.shape, dtype=bool) seeds_prev = np.zeros(im.shape, dtype=bool) max_radius = radii[0] if radii.size else 0 @@ -541,17 +477,30 @@ def local_thickness_bf( indices = _get_lt_flat_indices(seeds & ~seeds_prev) indices = _remove_lt_contained_disks(indices, seeds_prev) value = i + 1 if return_indices else radius - if indices.size: - lt = _insert_lt_disks_at_indices( - lt=lt, + if use_interface: + np.not_equal(lt, 0, out=nwp) + if indices.size: + nwp = _insert_disks_at_indices_parallel( + im=nwp, + indices=indices, + dt=dt, + ceil_distance=ceil_distance, + smooth=smooth, + overwrite=True, + fixed_radius=radius, + ) + nwp[seeds] = True + lt[(lt == 0) & nwp] = value + elif indices.size: + lt = _insert_disks_at_indices_parallel_direct( + im=lt, indices=indices, - radius=radius, - value=value, + dt=dt, ceil_distance=ceil_distance, smooth=smooth, + fixed_radius=radius, + value=value, ) - if use_interface: - lt[(lt == 0) & seeds] = value seeds_prev = seeds previous_radius = radius if return_indices: diff --git a/src/porespy/simulations/_tools.py b/src/porespy/simulations/_tools.py index 9efaf1dee..505316fc2 100644 --- a/src/porespy/simulations/_tools.py +++ b/src/porespy/simulations/_tools.py @@ -1,6 +1,11 @@ import numpy as np -from numba import get_num_threads, get_thread_id, njit, prange -from porespy.tools import _get_axial_extent +from numba import njit, prange +from porespy.tools import ( + _insert_disks_at_indices_parallel as _insert_disks_at_indices_parallel, + _insert_disks_at_indices_parallel_direct as _insert_disks_at_indices_parallel_direct, + _insert_disks_at_indices_parallel_merged as _insert_disks_at_indices_parallel_merged, + _use_merged_intervals as _use_merged_intervals, +) def _get_flat_indices(mask): @@ -89,245 +94,3 @@ def _find_interface(mask, interface): # pragma: no cover or (k + 1 < zlim and not mask[i, j, k + 1]) ) return interface - - -def _insert_disks_at_indices_parallel( - im, - indices, - dt, - ceil_distance, - smooth=True, - overwrite=False, -): # pragma: no cover - if overwrite and _use_merged_intervals(im, indices, dt): - return _insert_disks_at_indices_parallel_merged( - im=im, - indices=indices, - dt=dt, - ceil_distance=ceil_distance, - smooth=smooth, - ) - return _insert_disks_at_indices_parallel_direct( - im=im, - indices=indices, - dt=dt, - ceil_distance=ceil_distance, - smooth=smooth, - overwrite=overwrite, - ) - - -@njit -def _use_merged_intervals(im, indices, dt): - """Sample sphere sizes to choose between direct and merged scan-line writes.""" - if len(indices) == 0: - return False - nsamples = min(len(indices), 256) - estimated_intervals = 0 - if im.ndim == 2: - ylim = im.shape[1] - for q in range(nsamples): - ind = indices[q * len(indices) // nsamples] - i = ind // ylim - j = ind - i * ylim - estimated_intervals += 2 * int(dt[i, j]) + 1 - nrows = im.shape[0] - else: - ylim, zlim = im.shape[1:] - stride0 = ylim * zlim - for q in range(nsamples): - ind = indices[q * len(indices) // nsamples] - i = ind // stride0 - rem = ind - i * stride0 - j = rem // zlim - k = rem - j * zlim - diameter = 2 * int(dt[i, j, k]) + 1 - estimated_intervals += diameter**2 - nrows = im.shape[0] * im.shape[1] - # Row buffers and their initialization cost more than direct writes for - # sparse disks. Benchmarks place the crossover near 256 generated intervals - # per output row, while strongly overlapping spheres exceed this by orders - # of magnitude. - return estimated_intervals * len(indices) >= 256 * nrows * nsamples - - -@njit(parallel=True) -def _insert_disks_at_indices_parallel_direct( - im, - indices, - dt, - ceil_distance, - smooth=True, - overwrite=False, -): # pragma: no cover - npts = len(indices) - if im.ndim == 2: - xlim, ylim = im.shape - for q in prange(npts): - ind = indices[q] - i = ind // ylim - j = ind - i * ylim - r = int(dt[i, j]) - radius_squared = r**2 - for x in range(max(0, i - r), min(i + r + 1, xlim)): - dx = x - i - y_extent = _get_axial_extent( - radius_squared - dx**2, - ceil_distance, - smooth, - ) - y_start = max(0, j - y_extent) - y_stop = min(j + y_extent + 1, ylim) - if overwrite: - im[x, y_start:y_stop] = True - else: - for y in range(y_start, y_stop): - if not im[x, y]: - im[x, y] = True - elif im.ndim == 3: - xlim, ylim, zlim = im.shape - stride0 = ylim * zlim - for q in prange(npts): - ind = indices[q] - i = ind // stride0 - rem = ind - i * stride0 - j = rem // zlim - k = rem - j * zlim - r = int(dt[i, j, k]) - radius_squared = r**2 - for x in range(max(0, i - r), min(i + r + 1, xlim)): - dx = x - i - yz_extent = _get_axial_extent( - radius_squared - dx**2, - ceil_distance, - smooth, - ) - for y in range( - max(0, j - yz_extent), - min(j + yz_extent + 1, ylim), - ): - dy = y - j - z_extent = _get_axial_extent( - radius_squared - dx**2 - dy**2, - ceil_distance, - smooth, - ) - z_start = max(0, k - z_extent) - z_stop = min(k + z_extent + 1, zlim) - if overwrite: - im[x, y, z_start:z_stop] = True - else: - for z in range(z_start, z_stop): - if not im[x, y, z]: - im[x, y, z] = True - return im - - -@njit(parallel=True) -def _insert_disks_at_indices_parallel_merged( - im, - indices, - dt, - ceil_distance, - smooth=True, -): # pragma: no cover - """Insert disks by merging consecutive overlapping scan-line intervals.""" - nthreads = get_num_threads() - if im.ndim == 2: - xlim, ylim = im.shape - starts = np.full((nthreads, xlim), ylim, dtype=np.int32) - stops = np.zeros((nthreads, xlim), dtype=np.int32) - for q in prange(len(indices)): - thread = get_thread_id() - ind = indices[q] - i = ind // ylim - j = ind - i * ylim - r = int(dt[i, j]) - radius_squared = r**2 - for x in range(max(0, i - r), min(i + r + 1, xlim)): - dx = x - i - extent = _get_axial_extent( - radius_squared - dx**2, - ceil_distance, - smooth, - ) - start = max(0, j - extent) - stop = min(j + extent + 1, ylim) - if start >= stop: - continue - old_start = starts[thread, x] - old_stop = stops[thread, x] - if old_start == ylim: - starts[thread, x] = start - stops[thread, x] = stop - elif (start <= old_stop) and (stop >= old_start): - starts[thread, x] = min(start, old_start) - stops[thread, x] = max(stop, old_stop) - else: - im[x, old_start:old_stop] = True - starts[thread, x] = start - stops[thread, x] = stop - for q in prange(nthreads * xlim): - thread = q // xlim - x = q - thread * xlim - start = starts[thread, x] - if start < ylim: - im[x, start:stops[thread, x]] = True - elif im.ndim == 3: - xlim, ylim, zlim = im.shape - stride0 = ylim * zlim - nrows = xlim * ylim - starts = np.full((nthreads, nrows), zlim, dtype=np.int32) - stops = np.zeros((nthreads, nrows), dtype=np.int32) - for q in prange(len(indices)): - thread = get_thread_id() - ind = indices[q] - i = ind // stride0 - rem = ind - i * stride0 - j = rem // zlim - k = rem - j * zlim - r = int(dt[i, j, k]) - radius_squared = r**2 - for x in range(max(0, i - r), min(i + r + 1, xlim)): - dx = x - i - yz_extent = _get_axial_extent( - radius_squared - dx**2, - ceil_distance, - smooth, - ) - for y in range( - max(0, j - yz_extent), - min(j + yz_extent + 1, ylim), - ): - dy = y - j - z_extent = _get_axial_extent( - radius_squared - dx**2 - dy**2, - ceil_distance, - smooth, - ) - start = max(0, k - z_extent) - stop = min(k + z_extent + 1, zlim) - if start >= stop: - continue - row = x * ylim + y - old_start = starts[thread, row] - old_stop = stops[thread, row] - if old_start == zlim: - starts[thread, row] = start - stops[thread, row] = stop - elif (start <= old_stop) and (stop >= old_start): - starts[thread, row] = min(start, old_start) - stops[thread, row] = max(stop, old_stop) - else: - im[x, y, old_start:old_stop] = True - starts[thread, row] = start - stops[thread, row] = stop - for q in prange(nthreads * nrows): - thread = q // nrows - row = q - thread * nrows - start = starts[thread, row] - if start < zlim: - x = row // ylim - y = row - x * ylim - im[x, y, start:stops[thread, row]] = True - return im diff --git a/src/porespy/tools/_sphere_insertions.py b/src/porespy/tools/_sphere_insertions.py index ca0e15c96..c13acf6aa 100644 --- a/src/porespy/tools/_sphere_insertions.py +++ b/src/porespy/tools/_sphere_insertions.py @@ -1,5 +1,5 @@ import numpy as np -from numba import njit, prange +from numba import get_num_threads, get_thread_id, njit, prange __all__ = [ '_make_disk', @@ -8,6 +8,10 @@ '_make_balls', '_make_axial_extent_lookup', '_get_axial_extent', + '_insert_disks_at_indices_parallel', + '_insert_disks_at_indices_parallel_direct', + '_insert_disks_at_indices_parallel_merged', + '_use_merged_intervals', '_insert_disk_at_points', '_insert_disk_at_point', '_insert_disk_at_points_parallel', @@ -88,6 +92,258 @@ def _get_axial_extent(distance_squared, ceil_distance, smooth): return extent +def _insert_disks_at_indices_parallel( + im, + indices, + dt, + ceil_distance, + smooth=True, + overwrite=False, + fixed_radius=-1, +): # pragma: no cover + """Insert disks from flat indices using direct or merged scan-line writes.""" + if overwrite and _use_merged_intervals(im, indices, dt, fixed_radius): + return _insert_disks_at_indices_parallel_merged( + im=im, + indices=indices, + dt=dt, + ceil_distance=ceil_distance, + smooth=smooth, + fixed_radius=fixed_radius, + ) + return _insert_disks_at_indices_parallel_direct( + im=im, + indices=indices, + dt=dt, + ceil_distance=ceil_distance, + smooth=smooth, + overwrite=overwrite, + fixed_radius=fixed_radius, + ) + + +@njit +def _use_merged_intervals(im, indices, dt, fixed_radius=-1): + """Estimate whether merging scan-line intervals will reduce write work.""" + if len(indices) == 0: + return False + nsamples = min(len(indices), 256) + estimated_intervals = 0 + if im.ndim == 2: + ylim = im.shape[1] + for q in range(nsamples): + ind = indices[q * len(indices) // nsamples] + i = ind // ylim + j = ind - i * ylim + r = fixed_radius if fixed_radius >= 0 else int(dt[i, j]) + estimated_intervals += 2 * r + 1 + nrows = im.shape[0] + else: + ylim, zlim = im.shape[1:] + stride0 = ylim * zlim + for q in range(nsamples): + ind = indices[q * len(indices) // nsamples] + i = ind // stride0 + rem = ind - i * stride0 + j = rem // zlim + k = rem - j * zlim + r = fixed_radius if fixed_radius >= 0 else int(dt[i, j, k]) + diameter = 2 * r + 1 + estimated_intervals += diameter**2 + nrows = im.shape[0] * im.shape[1] + # Row buffers and their initialization cost more than direct writes for + # sparse disks. Benchmarks place the crossover near 256 generated intervals + # per output row, while strongly overlapping spheres exceed this by orders + # of magnitude. + return estimated_intervals * len(indices) >= 256 * nrows * nsamples + + +@njit(parallel=True) +def _insert_disks_at_indices_parallel_direct( + im, + indices, + dt, + ceil_distance, + smooth=True, + overwrite=False, + fixed_radius=-1, + value=True, +): # pragma: no cover + """Insert disks from flat indices using direct scan-line writes.""" + npts = len(indices) + if im.ndim == 2: + xlim, ylim = im.shape + for q in prange(npts): + ind = indices[q] + i = ind // ylim + j = ind - i * ylim + r = fixed_radius if fixed_radius >= 0 else int(dt[i, j]) + radius_squared = r**2 + for x in range(max(0, i - r), min(i + r + 1, xlim)): + dx = x - i + y_extent = _get_axial_extent( + radius_squared - dx**2, + ceil_distance, + smooth, + ) + y_start = max(0, j - y_extent) + y_stop = min(j + y_extent + 1, ylim) + if overwrite: + im[x, y_start:y_stop] = value + else: + for y in range(y_start, y_stop): + if not im[x, y]: + im[x, y] = value + elif im.ndim == 3: + xlim, ylim, zlim = im.shape + stride0 = ylim * zlim + for q in prange(npts): + ind = indices[q] + i = ind // stride0 + rem = ind - i * stride0 + j = rem // zlim + k = rem - j * zlim + r = fixed_radius if fixed_radius >= 0 else int(dt[i, j, k]) + radius_squared = r**2 + for x in range(max(0, i - r), min(i + r + 1, xlim)): + dx = x - i + yz_extent = _get_axial_extent( + radius_squared - dx**2, + ceil_distance, + smooth, + ) + for y in range( + max(0, j - yz_extent), + min(j + yz_extent + 1, ylim), + ): + dy = y - j + z_extent = _get_axial_extent( + radius_squared - dx**2 - dy**2, + ceil_distance, + smooth, + ) + z_start = max(0, k - z_extent) + z_stop = min(k + z_extent + 1, zlim) + if overwrite: + im[x, y, z_start:z_stop] = value + else: + for z in range(z_start, z_stop): + if not im[x, y, z]: + im[x, y, z] = value + return im + + +@njit(parallel=True) +def _insert_disks_at_indices_parallel_merged( + im, + indices, + dt, + ceil_distance, + smooth=True, + fixed_radius=-1, +): # pragma: no cover + """Insert disks by merging consecutive overlapping scan-line intervals.""" + nthreads = get_num_threads() + if im.ndim == 2: + xlim, ylim = im.shape + starts = np.full((nthreads, xlim), ylim, dtype=np.int32) + stops = np.zeros((nthreads, xlim), dtype=np.int32) + for q in prange(len(indices)): + thread = get_thread_id() + ind = indices[q] + i = ind // ylim + j = ind - i * ylim + r = fixed_radius if fixed_radius >= 0 else int(dt[i, j]) + radius_squared = r**2 + for x in range(max(0, i - r), min(i + r + 1, xlim)): + dx = x - i + extent = _get_axial_extent( + radius_squared - dx**2, + ceil_distance, + smooth, + ) + start = max(0, j - extent) + stop = min(j + extent + 1, ylim) + if start >= stop: + continue + old_start = starts[thread, x] + old_stop = stops[thread, x] + if old_start == ylim: + starts[thread, x] = start + stops[thread, x] = stop + elif (start <= old_stop) and (stop >= old_start): + starts[thread, x] = min(start, old_start) + stops[thread, x] = max(stop, old_stop) + else: + im[x, old_start:old_stop] = True + starts[thread, x] = start + stops[thread, x] = stop + for q in prange(nthreads * xlim): + thread = q // xlim + x = q - thread * xlim + start = starts[thread, x] + if start < ylim: + im[x, start:stops[thread, x]] = True + elif im.ndim == 3: + xlim, ylim, zlim = im.shape + stride0 = ylim * zlim + nrows = xlim * ylim + starts = np.full((nthreads, nrows), zlim, dtype=np.int32) + stops = np.zeros((nthreads, nrows), dtype=np.int32) + for q in prange(len(indices)): + thread = get_thread_id() + ind = indices[q] + i = ind // stride0 + rem = ind - i * stride0 + j = rem // zlim + k = rem - j * zlim + r = fixed_radius if fixed_radius >= 0 else int(dt[i, j, k]) + radius_squared = r**2 + for x in range(max(0, i - r), min(i + r + 1, xlim)): + dx = x - i + yz_extent = _get_axial_extent( + radius_squared - dx**2, + ceil_distance, + smooth, + ) + for y in range( + max(0, j - yz_extent), + min(j + yz_extent + 1, ylim), + ): + dy = y - j + z_extent = _get_axial_extent( + radius_squared - dx**2 - dy**2, + ceil_distance, + smooth, + ) + start = max(0, k - z_extent) + stop = min(k + z_extent + 1, zlim) + if start >= stop: + continue + row = x * ylim + y + old_start = starts[thread, row] + old_stop = stops[thread, row] + if old_start == zlim: + starts[thread, row] = start + stops[thread, row] = stop + elif (start <= old_stop) and (stop >= old_start): + starts[thread, row] = min(start, old_start) + stops[thread, row] = max(stop, old_stop) + else: + im[x, y, old_start:old_stop] = True + starts[thread, row] = start + stops[thread, row] = stop + for q in prange(nthreads * nrows): + thread = q // nrows + row = q - thread * nrows + start = starts[thread, row] + if start < zlim: + x = row // ylim + y = row - x * ylim + im[x, y, start:stops[thread, row]] = True + return im + + @njit(parallel=True) def _insert_disks_at_points_parallel(im, coords, radii, v, smooth=True, overwrite=False): # pragma: no cover diff --git a/test/unit/test_simulations_ibop.py b/test/unit/test_simulations_ibop.py index 027d51cfc..e8b18322d 100644 --- a/test/unit/test_simulations_ibop.py +++ b/test/unit/test_simulations_ibop.py @@ -132,6 +132,46 @@ def test_merged_interval_sphere_insertion(self): ) assert np.array_equal(adaptive, direct) + def test_flat_index_insertion_with_fixed_radius_and_value(self): + for shape in [(31, 37), (19, 23, 17)]: + centers = np.zeros(shape, dtype=bool) + center_slice = tuple(slice(10, 20) for _ in shape) + centers[center_slice] = True + indices = _get_flat_indices(centers) + coords = np.vstack(np.where(centers)) + dt = np.zeros(shape, dtype=np.float32) + dt[centers] = 8 + radius = 3 + lookup = _make_axial_extent_lookup(radius) + expected = ps.tools._insert_disks_at_points_parallel( + im=np.zeros(shape, dtype=bool), + coords=coords, + radii=np.full(coords.shape[1], radius), + v=True, + smooth=True, + ) + actual = _insert_disks_at_indices_parallel( + im=np.zeros(shape, dtype=bool), + indices=indices, + dt=dt, + ceil_distance=lookup, + smooth=True, + overwrite=True, + fixed_radius=radius, + ) + assert np.array_equal(actual, expected) + + labels = _insert_disks_at_indices_parallel_direct( + im=np.zeros(shape, dtype=np.uint8), + indices=indices, + dt=dt, + ceil_distance=lookup, + smooth=True, + fixed_radius=radius, + value=7, + ) + assert np.array_equal(labels, expected * 7) + def test_interval_merging_is_adaptive(self): shape = (51, 53, 49) dt = np.zeros(shape, dtype=np.float32)