如何用不同维度的mask替换嵌套numpy数组中的元素?
解决Numpy嵌套数组按掩码替换子数组的问题
错误原因
你的frm转成Numpy数组后shape为(3, 2, 2),是三维数组(3组,每组包含2个长度为2的子数组),而mask是(3, 2)的二维布尔数组。直接使用to[mask]时,Numpy会错误地将掩码视为对数组第一个维度的索引,导致布尔索引长度与数组维度不匹配,触发IndexError。同时即使维度匹配,直接赋值也无法对应到子数组的维度。
修正方案
方案1:扩展掩码维度实现广播赋值
给掩码增加一个维度,让它与数组的最后一维匹配,通过广播机制实现对子数组的批量替换:
import numpy as np frm = [[[1,2], [2,3]], [[3,4], [4,5]], [[5,6], [6,7]]] mask = [[False, True], [False, True], [True, False]] repl = [0,0] frm = np.array(frm) to = frm.copy() # 给掩码增加最后一个维度,使其shape变为(3,2,1),与数组前两维匹配 to[np.expand_dims(mask, axis=-1)] = repl print(to)
输出结果:
[[[1 2] [0 0]] [[3 4] [0 0]] [[0 0] [6 7]]]
方案2:使用np.where定位索引赋值
通过np.where获取掩码中True元素的行列索引,直接定位到需要替换的子数组进行赋值,逻辑更直观:
import numpy as np frm = [[[1,2], [2,3]], [[3,4], [4,5]], [[5,6], [6,7]]] mask = [[False, True], [False, True], [True, False]] repl = [0,0] frm = np.array(frm) to = frm.copy() # 获取掩码为True的位置的行、列索引 rows, cols = np.where(mask) # 对指定位置的子数组赋值 to[rows, cols] = repl print(to)
输出结果与方案1完全一致。
内容的提问来源于stack exchange,提问作者Horst
相关产品推荐
相关产品推荐

