向量化改进2D数组严格比较下的局部极值查找函数
优化2D NumPy数组局部极值检测函数的向量化方案
需求说明
需要实现高效函数,返回2D NumPy数组的严格局部极小值/极大值:元素需严格小于/大于滑动窗口内的所有邻居(窗口大小为奇数且≥3)。skimage.morphology.local_minimum/local_maxima采用非严格比较(≤/≥),不符合需求。
现有实现分析
1. 滑动窗口循环实现(性能瓶颈)
该方案通过sliding_window_view生成窗口后逐一遍历,功能正确但因Python循环导致性能低下:
import numpy as np def get_local_extrema(array, window_size=(3, 3)): if not all(size % 2 == 1 and size >= 3 for size in window_size): raise ValueError("Window size must be odd and >= 3 in both dimensions.") minima_map = np.zeros_like(array) maxima_map = np.zeros_like(array) original_size = array.shape half_window_size = tuple(size // 2 for size in window_size) padded_array = np.pad(array.astype(float), tuple((size, size) for size in half_window_size), mode='constant', constant_values=np.nan) windows = np.lib.stride_tricks.sliding_window_view(padded_array, window_size).reshape( original_size[0] * original_size[1], *window_size) mask = np.ones(window_size, dtype=bool) mask[half_window_size] = False for i in range(windows.shape[0]): window = windows[i] center_val = window[half_window_size] masked_window = window[mask] row = i // original_size[1] col = i % original_size[1] if center_val > np.nanmax(masked_window): maxima_map[row, col] = center_val elif center_val < np.nanmin(masked_window): minima_map[row, col] = center_val return minima_map.astype(array.dtype), maxima_map.astype(array.dtype)
2. 手动向量化实现(边界缺失)
该方案通过数组切片实现向量化,但仅处理内部元素,无法检测边界上的极值:
def get_local_extrema_2(img): minima_map = np.zeros_like(img) maxima_map = np.zeros_like(img) minima_map[1:-1, 1:-1] = np.where( (img[1:-1, 1:-1] < img[:-2, 1:-1]) & (img[1:-1, 1:-1] < img[2:, 1:-1]) & (img[1:-1, 1:-1] < img[1:-1, :-2]) & (img[1:-1, 1:-1] < img[1:-1, 2:]) & (img[1:-1, 1:-1] < img[2:, 2:]) & (img[1:-1, 1:-1] < img[:-2, :-2]) & (img[1:-1, 1:-1] < img[2:, :-2]) & (img[1:-1, 1:-1] < img[:-2, 2:]), img[1:-1, 1:-1], 0) maxima_map[1:-1, 1:-1] = np.where( (img[1:-1, 1:-1] > img[:-2, 1:-1]) & (img[1:-1, 1:-1] > img[2:, 1:-1]) & (img[1:-1, 1:-1] > img[1:-1, :-2]) & (img[1:-1, 1:-1] > img[1:-1, 2:]) & (img[1:-1, 1:-1] > img[2:, 2:]) & (img[1:-1, 1:-1] > img[:-2, :-2]) & (img[1:-1, 1:-1] > img[2:, :-2]) & (img[1:-1, 1:-1] > img[:-2, 2:]), img[1:-1, 1:-1], 0) return minima_map, maxima_map
3. 基于SciPy形态学的高效实现
该方案利用scipy.ndimage的腐蚀/膨胀操作(C级实现),高效处理边界且符合严格比较要求:
import numpy as np import scipy def get_local_extrema_v3(image): footprint = np.ones((3, 3), dtype=bool) footprint[1, 1] = False # 腐蚀操作取窗口内邻居的最小值,若最小值>中心则为严格极小 minima = image * (scipy.ndimage.grey_erosion(image, footprint=footprint, mode='mirror') > image) # 膨胀操作取窗口内邻居的最大值,若最大值<中心则为严格极大 maxima = image * (scipy.ndimage.grey_dilation(image, footprint=footprint, mode='mirror') < image) return minima, maxima
更多向量化优化建议
1. 支持任意奇数窗口大小(SciPy扩展)
修改footprint适配任意奇数窗口,同时保留边界处理能力:
def get_local_extrema_scipy_variable(image, window_size=(3,3)): if not all(size %2 ==1 and size>=3 for size in window_size): raise ValueError("Window size must be odd and >=3") half_h, half_w = [s//2 for s in window_size] footprint = np.ones(window_size, dtype=bool) footprint[half_h, half_w] = False # 可选mode:mirror/constant/nearest等,根据边界需求选择 erosion = scipy.ndimage.grey_erosion(image, footprint=footprint, mode='mirror') dilation = scipy.ndimage.grey_dilation(image, footprint=footprint, mode='mirror') minima = image * (erosion > image) maxima = image * (dilation < image) return minima, maxima
2. 纯NumPy向量化实现(无SciPy依赖)
利用sliding_window_view结合广播操作,完全避免Python循环,支持任意窗口:
import numpy as np def get_local_extrema_numpy(array, window_size=(3,3)): if not all(size %2 ==1 and size>=3 for size in window_size): raise ValueError("Window size must be odd and >=3") half_h, half_w = [s//2 for s in window_size] # 边界填充inf,确保边界元素的邻居比较逻辑正确 pad_width = ((half_h, half_h), (half_w, half_w)) padded = np.pad(array.astype(np.float64), pad_width, mode='constant', constant_values=np.inf) # 生成滑动窗口 windows = np.lib.stride_tricks.sliding_window_view(padded, window_size) # 提取中心元素并扩展维度,用于广播比较 center = array[..., np.newaxis, np.newaxis] # 生成掩码排除中心元素 mask = np.ones(window_size, dtype=bool) mask[half_h, half_w] = False # 向量化判断所有窗口的极值条件 is_min = np.all(windows[:, :, mask] > center, axis=(2,3)) is_max = np.all(windows[:, :, mask] < center, axis=(2,3)) # 转换回原数组类型 minima_map = np.where(is_min, array, 0).astype(array.dtype) maxima_map = np.where(is_max, array, 0).astype(array.dtype) return minima_map, maxima_map
3. 性能优化细节
- 类型精简:避免不必要的浮点转换,比如用
np.iinfo(array.dtype).max/min代替np.inf处理整数数组的边界填充,减少内存开销。 - 边界模式选择:根据业务需求选择
pad或ndimage的mode参数:constant适合将边界外视为极值,mirror适合镜像延伸数据,nearest适合最近值填充。 - 布尔值输出:若无需保留原极值数值,可直接返回布尔掩码,进一步提升性能:
is_min和is_max本身就是布尔数组,无需额外转换。
4. 极端大数组优化
对于超大数组,可结合numba的JIT编译进一步加速纯NumPy逻辑,或分块处理数组减少内存占用:
from numba import jit @jit(nopython=True) def get_local_extrema_numba(array, window_size=(3,3)): rows, cols = array.shape half_h, half_w = window_size[0]//2, window_size[1]//2 minima_map = np.zeros_like(array) maxima_map = np.zeros_like(array) for i in range(rows): for j in range(cols): # 手动计算窗口范围,处理边界 row_start = max(0, i - half_h) row_end = min(rows, i + half_h + 1) col_start = max(0, j - half_w) col_end = min(cols, j + half_w + 1) window = array[row_start:row_end, col_start:col_end] center_val = array[i,j] # 排除中心元素 window = window.flatten() window = window[window != center_val] if np.all(window > center_val): minima_map[i,j] = center_val elif np.all(window < center_val): maxima_map[i,j] = center_val return minima_map, maxima_map
内容的提问来源于stack exchange,提问作者Oliver
相关产品推荐
相关产品推荐

