You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

向量化改进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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.06 03:00:56