如何在NumPy中使用搜索窗口保留最大值并将其余值置零
实现滑动窗口保留最大值并置零其他元素(NumPy版)
嘿,我完全get到你的需求了——让一个指定大小的滑动窗口遍历二维数组,每个窗口里只留下最大值,其他所有元素通通置零。我用NumPy给你写个实用的实现,还会结合你给的小例子来验证效果,保证清晰易懂!
核心思路
咱先理清楚要做的事:
- 用NumPy的滑动窗口工具生成所有窗口视图(高效不占额外内存);
- 找出每个窗口的最大值,以及它在窗口内的位置;
- 把最大值映射回原数组的对应位置,其他位置全部设为0。
针对你示例的代码实现
先拿你给的小数组测试,窗口用你提到的2×2:
import numpy as np # 你的示例数组 x = np.array([[1,2,3,4,5,6,7,8,9,10], [2,5,4,5,3,4,6,7,5,3], [3,3,4,5,6,7,3,4,5,8]]) def sliding_window_keep_max(arr, window_size): arr_h, arr_w = arr.shape win_h, win_w = window_size # 生成滑动窗口视图,形状为 (窗口行数, 窗口列数, 窗口高, 窗口宽) windows = np.lib.stride_tricks.sliding_window_view(arr, window_size) # 获取每个窗口的最大值,以及最大值在窗口内的一维索引 max_vals = np.max(windows, axis=(2, 3)) max_idx = np.argmax(windows.reshape(windows.shape[0], windows.shape[1], -1), axis=2) # 把一维索引转成窗口内的二维坐标 max_win_h = max_idx // win_w max_win_w = max_idx % win_w # 初始化全零结果数组 result = np.zeros_like(arr) # 遍历每个窗口的起始位置 for i in range(windows.shape[0]): for j in range(windows.shape[1]): # 映射到原数组的坐标 orig_h = i + max_win_h[i, j] orig_w = j + max_win_w[i, j] # 保留最大值 result[orig_h, orig_w] = max_vals[i, j] return result # 处理示例数组 result = sliding_window_keep_max(x, (2, 2)) print("处理后的结果:") print(result)
运行这段代码后,第一个2×2窗口(左上角的[[1,2],[2,5]])里的最大值是5,它在原数组的位置是(1,1),所以这个位置保留5,窗口内其他位置(0,0)、(0,1)、(1,0)都被置零,完全符合你描述的Result = [[0,0...],[0,5...]]的效果!
扩展:支持步长和多最大值保留
如果你的需求有变化,比如想要更大的窗口(比如3×3)、自定义滑动步长,或者窗口内有多个相同最大值时要全部保留,可以用这个增强版函数:
def sliding_window_keep_max_enhanced(arr, window_size, stride=(1, 1), keep_all_max=True): arr_h, arr_w = arr.shape win_h, win_w = window_size stride_h, stride_w = stride # 计算能容纳的窗口数量 num_win_h = (arr_h - win_h) // stride_h + 1 num_win_w = (arr_w - win_w) // stride_w + 1 result = np.zeros_like(arr) for i in range(num_win_h): for j in range(num_win_w): # 提取当前窗口 win_start_h = i * stride_h win_start_w = j * stride_w window = arr[win_start_h:win_start_h+win_h, win_start_w:win_start_w+win_w] max_val = np.max(window) if keep_all_max: # 找到窗口内所有最大值的位置 max_pos = np.where(window == max_val) orig_h = win_start_h + max_pos[0] orig_w = win_start_w + max_pos[1] result[orig_h, orig_w] = max_val else: # 只保留第一个最大值 max_idx = np.argmax(window) max_win_h = max_idx // win_w max_win_w = max_idx % win_w result[win_start_h + max_win_h, win_start_w + max_win_w] = max_val return result # 测试3×3窗口、步长1的情况(适配你提到的101行100列数组) # big_arr = np.random.rand(101, 100) # big_result = sliding_window_keep_max_enhanced(big_arr, (3, 3))
这个版本更灵活,不管是你说的101×100的大数组,还是其他需求场景,都能轻松应对。
内容的提问来源于stack exchange,提问作者Simba Kapfumo
相关产品推荐
相关产品推荐

