如何基于epoch_label掩码对numpy二维数组切片获取目标行
问题原因
- 你使用
ma.masked_where时逻辑搞反:ma.masked_where(条件, 数组)会标记数组中满足条件的位置为掩码,你设置的条件为epoch_label == 0,相当于把你需要保留的行全部标记为无效,剩下的有效位置反而是标签为1的行,因此索引结果完全错误。 - 首次使用
np.where的逻辑错误:仅传入单个数组时np.where会返回数组中非零/未掩码元素的坐标,不会直接返回筛选后的数组内容。 - 第二次直接用掩码数组索引时,数组会读取掩码中未被标记的数值作为索引值,你的示例中未被标记的是前两个位置的数值1、1,等价于你连续取了两次索引1的行,和你预期完全不符。
正确解决方案
方案1:直接使用布尔索引(最简便,不需要引入掩码数组)
这是numpy筛选符合条件行的标准用法,代码如下:
# 生成布尔掩码,True对应需要保留的行(epoch_label等于0的位置) mask = epoch_label == 0 expected_output = epoch_com_arr[mask]
运行后即可直接得到你需要的[[2 4 7 6],[8 8 1 6]]结果。
方案2:修正掩码数组的用法
如果你需要保留masked_where的写法,需要调整条件和索引逻辑:
# 条件反过来:mask掉不需要的行(epoch_label等于1的位置) mm = ma.masked_where(epoch_label == 1, epoch_label) # 取掩码的反选作为索引条件,True为需要保留的行 expected_output = epoch_com_arr[~mm.mask]
内容的提问来源于stack exchange,提问作者rpb
相关产品推荐
相关产品推荐

