如何使用np.where处理三维数组,匹配并替换指定子数组?
匹配并修改三维数组中的完整子数组
给定如下三维NumPy数组:
import numpy as np a3 = np.array([ [[1, 2, 3], [4, 5,6], [1, 8, 9]], [[1, 2, 3], [4, 5, 6], [1, 8, 9]], [[1, 2, 3], [4, 5, 6], [1, 8, 9]] ])
需求是找到所有值为[1,2,3]的完整子数组,将其替换为[9,9,9],其余子数组替换为[0,0,0]。
直接使用np.where(a3 == [1,2,3], [9,9,9], [0,0,0])会得到不符合预期的结果,因为该操作是逐元素匹配,而非判断整个子数组是否完全一致——比如第三个子数组的第一个元素是1,会被单独替换为9,导致出现[9,0,0]这样的错误结果。
解决方案
要实现完整子数组的匹配,需要先生成标记每个子数组是否完全匹配的布尔掩码,再基于掩码进行替换:
方法1:先创建掩码再赋值
- 生成二维掩码,标记每个子数组是否完全等于
[1,2,3]:
# 沿着最后一个轴(子数组的元素轴)判断所有元素是否匹配 mask = np.all(a3 == [1,2,3], axis=2)
此时mask的形状为(3,3),每个True对应原数组中一个完全匹配的子数组。
- 初始化结果数组并替换目标子数组:
result = np.zeros_like(a3) # 将掩码为True的位置替换为[9,9,9] result[mask] = [9,9,9]
方法2:结合np.expand_dims和np.where
将二维掩码扩展为三维(与原数组形状一致),直接用np.where完成替换:
mask_3d = np.expand_dims(mask, axis=2) result = np.where(mask_3d, [9,9,9], [0,0,0])
两种方法最终都会得到期望的结果:
[[[9 9 9] [0 0 0] [0 0 0]] [[9 9 9] [0 0 0] [0 0 0]] [[9 9 9] [0 0 0] [0 0 0]]]
内容的提问来源于stack exchange,提问作者kerz
相关产品推荐
相关产品推荐

