如何用2D掩码高效过滤3D数组?numpy替代循环方案
用Numpy掩码高效筛选多维数组
我有一个形状为(m,n,3)的data数组,想通过形状为(m,n)的掩码对其值进行筛选,最终得到形状为(x,3)的output数组。下面的循环代码能实现需求,但想替换掉循环来获得更高效的实现方案:
import numpy as np data = np.array([ [[11, 12, 13], [14, 15, 16], [17, 18, 19]], [[21, 22, 13], [24, 25, 26], [27, 28, 29]], [[31, 32, 33], [34, 35, 36], [37, 38, 39]], ]) mask = np.array([ [False, False, True], [False, True, False], [True, True, False], ]) output = [] for i in range(len(mask)): for j in range(len(mask[i])): if mask[i][j] == True: output.append(data[i][j]) output = np.array(output)
预期输出为:
np.array([[17, 18, 19], [24, 25, 26], [31, 32, 33], [34, 35, 36]])
高效解决方案
直接利用Numpy的布尔索引特性就能完成这个操作,无需嵌套循环,核心代码仅需一行:
output = data[mask]
原理说明
Numpy的布尔索引会自动识别掩码中True的位置,提取data中对应位置的元素,并自动将结果重塑为(x,3)的形状——其中x就是掩码里True的总数量。
验证结果
运行上述代码后,得到的结果和预期完全一致:
array([[17, 18, 19], [24, 25, 26], [31, 32, 33], [34, 35, 36]])
优势
这种方法基于Numpy底层的C语言优化实现,相比Python嵌套循环,在数组规模较大时性能提升极其显著,同时代码更简洁易读。
内容的提问来源于stack exchange,提问作者Florian Ludewig
相关产品推荐
相关产品推荐

