You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用不同维度的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 04:06:25