如何在NumPy数组中查找指定大小、数值在[μ±σ]区间的patch
实现思路
你可以借助NumPy的滑动窗口视图能力高效实现,不需要手动写嵌套循环遍历,核心逻辑如下:
- 首先生成布尔掩码:判断原数组每个元素是否落在
[mu - sigma, mu + sigma]区间内,符合条件标记为True - 调用
np.lib.stride_tricks.sliding_window_view切出所有尺寸为patch_size × patch_size的候选窗口,该方法不会额外复制数组,性能开销很低 - 对每个窗口执行
all()判断,只要窗口内所有值都是True,就说明这个patch完全符合要求 - 最后用
np.where/np.argwhere提取符合条件的窗口左上角坐标,按需转换为1基/0基索引即可
完整实现代码
import numpy as np from numpy.lib.stride_tricks import sliding_window_view a = np.array([[2, 1, 6, 7, 6, 5, 9, 1, 5, 6], [1, 7, 6, 0, 1, 9, 8, 1, 2, 0], [4, 4, 5, 1, 7, 8, 8, 7, 3, 3], [5, 6, 4, 4, 5, 4, 2, 2, 2, 7], [3, 4, 4, 5, 5, 4, 8, 6, 1, 9], [4, 4, 5, 5, 4, 6, 1, 9, 4, 5], [8, 4, 6, 4, 4, 5, 2, 1, 8, 0], [4, 5, 5, 5, 5, 4, 6, 2, 2, 4], [3, 6, 1, 7, 7, 3, 2, 3, 5, 1], [5, 1, 8, 3, 1, 4, 5, 9, 5, 0]]) patch_mu = 5 patch_sigma = 1 patch_size = 5 # 正方形边长直接传数值即可 def find_patch_index(arr, mu, sigma, size): # 边界合法性判断 if size > arr.shape[0] or size > arr.shape[1]: raise ValueError("patch尺寸不能大于原数组尺寸") # 生成区间匹配布尔掩码 lower = mu - sigma upper = mu + sigma mask = (arr >= lower) & (arr <= upper) # 切出所有指定尺寸的滑动窗口 windows = sliding_window_view(mask, window_shape=(size, size)) # 判断窗口内所有元素是否符合要求,返回0基索引 valid = windows.all(axis=(-1, -2)) idx = np.argwhere(valid) # 转换为1基索引和示例输出匹配 idx = idx + 1 # 可按需返回第一个符合条件的坐标,或所有符合条件的坐标列表 return idx[0] if len(idx) > 0 else None idx = find_patch_index(a, patch_mu, patch_sigma, patch_size) print(idx) # 输出 [3 1],和示例要求一致
如果你使用的NumPy版本低于1.20.0没有内置sliding_window_view,可以替换为np.lib.stride_tricks.as_strided实现滑动窗口,核心判断逻辑完全一致。
内容的提问来源于stack exchange,提问作者ravi
相关产品推荐
相关产品推荐

