如何高效实现矩阵行非邻域次大值的索引/窗口查找
问题:寻找避开最大值邻域的每行次大值索引
需求:给定二维数组,需先获取每行最大值的索引,再查找不在该索引±n(示例中n=2)范围内的每行次大值索引。
示例矩阵与当前实现
初始矩阵与最大值索引获取
import numpy as np results = np.array([ [ 33, 108, 208, 96, 96, 112, 18, 208, 33, 323, 60, 42], [ 51, 6, 39, 112, 160, 144, 342, 195, 27, 136, 42, 54], [ 12, 176, 266, 162, 45, 70, 156, 198, 143, 56, 342, 130], [ 22, 288, 304, 162, 21, 238, 156, 126, 165, 91, 144, 130], [342, 120, 36, 51, 10, 128, 156, 272, 32, 98, 192, 288] ]) # 获取每行最大值的索引 row_max_index = results.argmax(1) print(row_max_index) # 输出:array([ 9, 6, 10, 2, 0])
当前繁琐的实现方式
当前通过构造掩码将最大值邻域元素置0后再求次大值索引,代码如下:
n = 2 col_count = results.shape[1] # 构造最大值±n范围内的索引,取模处理边界 maskIndx = np.c_[ row_max_index - n, row_max_index - n + 1, row_max_index, row_max_index + 1, row_max_index + n ] % col_count print(maskIndx) # 输出: # array([[ 7, 8, 9, 10, 11], # [ 4, 5, 6, 7, 8], # [ 8, 9, 10, 11, 0], # [ 0, 1, 2, 3, 4], # [10, 11, 0, 1, 2]]) # 将邻域元素置0 results[np.meshgrid(np.arange(results.shape[0]), np.arange(2*n+1))[1], maskIndx] = 0 print(results) # 输出: # array([[ 33, 108, 208, 96, 96, 112, 18, 0, 0, 0, 0, 0], # [ 51, 6, 39, 112, 0, 0, 0, 0, 0, 136, 42, 54], # [ 0, 176, 266, 162, 45, 70, 156, 198, 0, 0, 0, 0], # [ 0, 0, 0, 0, 0, 238, 156, 126, 165, 91, 144, 130], # [ 0, 0, 0, 51, 10, 128, 156, 272, 32, 98, 0, 0]]) # 获取次大值索引 next_max_index = results.argmax(1) print(next_max_index) # 输出:array([2, 9, 2, 5, 7])
更高效的实现方案
不需要修改原数组,直接通过索引掩码筛选符合条件的元素,避免了原数组修改带来的开销,逻辑也更简洁:
方案1:广播构造掩码筛选
n = 2 row_count, col_count = results.shape # 构造每行的列索引矩阵 cols = np.arange(col_count)[np.newaxis, :].repeat(row_count, axis=0) # 计算每个列索引与最大值索引的距离(支持环形边界,若不需要环形则去掉minumum和取模部分) distance = np.minimum( np.abs(cols - row_max_index[:, np.newaxis]), col_count - np.abs(cols - row_max_index[:, np.newaxis]) ) # 不需要环形的话,直接用: # distance = np.abs(cols - row_max_index[:, np.newaxis]) # 将邻域内的元素设为极小值,不参与最大值计算 masked_results = np.where(distance > n, results, -np.inf) # 获取次大值索引 next_max_index = masked_results.argmax(1) print(next_max_index) # 输出:array([2, 9, 2, 5, 7])
方案2:排序后筛选有效索引
如果需要处理多个候选值,可先对每行元素按值降序排序索引,再过滤邻域内的索引:
n = 2 col_count = results.shape[1] # 获取每行元素按值降序排列的索引 sorted_indices = np.argsort(-results, axis=1) # 计算排序后索引与最大值索引的距离(支持环形边界) distance = np.minimum( np.abs(sorted_indices - row_max_index[:, np.newaxis]), col_count - np.abs(sorted_indices - row_max_index[:, np.newaxis]) ) # 找到每行第一个符合距离要求的索引 next_max_index = sorted_indices[np.arange(results.shape[0]), (distance > n).argmax(axis=1)] print(next_max_index) # 输出:array([2, 9, 2, 5, 7])
内容的提问来源于stack exchange,提问作者bigwillis234
相关产品推荐
相关产品推荐

