如何从numpy.where()输出生成布尔掩码并保留数组维度?
解决方法
1. 直接生成布尔掩码
你完全不需要通过np.where()来生成布尔掩码,条件表达式本身就是和原数组形状一致的布尔数组:
import numpy as np A = np.random.random((10, 10)) mask = A > 0.9 # 这就是(10,10)的布尔掩码,直接用于逐行处理
比如遍历每行掩码做后续操作:
for row_mask in mask: # 处理当前行的布尔标记 pass
2. 保留维度的筛选值数组
如果想要得到和原数组形状一致、符合条件的位置保留原值、其余位置设为NaN(或其他填充值),有几种更简洁的写法:
方法一:简化你的现有方案
不用单独创建全NaN数组,np.where()支持直接传入np.nan,它会自动广播成和A相同的形状:
C = np.where(A > 0.9, A, np.nan)
方法二:布尔掩码直接赋值
先创建一个填充好默认值的数组,再把符合条件的值填充进去,逻辑更直观:
C = np.full_like(A, np.nan) C[A > 0.9] = A[A > 0.9]
方法三:掩码数组(可选)
如果只是需要标记无效值而非替换,可以用np.ma.masked_array,它会保留原数组形状,同时标记不符合条件的位置:
masked_A = np.ma.masked_array(A, mask=~(A > 0.9))
后续可以通过masked_A.mask获取布尔掩码,masked_A.data获取原数据,同样支持逐行操作。
补充:为什么原写法会扁平化?
np.where()返回的是满足条件元素的索引元组(二维数组会返回(row_indices, col_indices)),用这个元组索引原数组时,numpy的高级索引特性会返回扁平化的一维数组。如果要保留形状,就得用布尔掩码直接操作,或者用np.where的三参数形式。
内容的提问来源于stack exchange,提问作者userE
相关产品推荐
相关产品推荐

