如何根据条件用numpy数组替换元素以生成三维数组?
问题:基于numpy条件替换生成三维数组
现有代码:
import numpy as np subst1 = np.array([2, 2, 2, 2]) subst2 = np.array([3, 3, 3, 3]) a = np.array([[1, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0,]]) b = np.where(0==a, subst1, subst2)
运行结果:
>>> a array([[1, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]) >>> b array([[3, 2, 2, 2], [2, 2, 2, 2], [2, 2, 2, 2]])
期望结果:
array([[[3,3,3,3], [2,2,2,2], [2,2,2,2], [2,2,2,2]], [[2,2,2,2], [2,2,2,2], [2,2,2,2], [2,2,2,2]], [[2,2,2,2], [2,2,2,2], [2,2,2,2], [2,2,2,2]]])
当前numpy.where的问题在于它是逐元素匹配替换,无法直接将原二维数组的每个元素替换为整个一维的subst数组。以下是两种高效的numpy原生解决方案:
方案1:利用广播扩展维度实现条件选择
通过将原数组a扩展为三维,使其维度与subst数组匹配,再结合numpy.where的广播特性完成替换,无需额外复制数据,性能最优:
import numpy as np subst1 = np.array([2, 2, 2, 2]) subst2 = np.array([3, 3, 3, 3]) a = np.array([[1, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]) # 将a扩展为三维(添加最后一个维度),shape变为(3,4,1) mask = a[..., np.newaxis] # 利用广播自动将subst1/subst2扩展为(3,4,4),完成条件替换 result = np.where(mask, subst2, subst1)
方案2:预填充数组后批量替换
先创建一个全为subst1的三维数组,再通过条件索引批量替换为subst2,逻辑直观:
import numpy as np subst1 = np.array([2, 2, 2, 2]) subst2 = np.array([3, 3, 3, 3]) a = np.array([[1, 0, 0, 0], [0, 0, 0, 0], [0, 0, 0, 0]]) # 生成与目标shape一致的三维数组,初始填充subst1 result = np.tile(subst1, (a.shape[0], a.shape[1], 1)) # 定位a中值为1的位置,批量替换为subst2 result[a == 1] = subst2
两种方案均能生成符合期望的三维数组,其中方案1借助numpy广播机制,在大数组场景下性能更优,适合用于性能对比测试。
内容的提问来源于stack exchange,提问作者Zoltan K.
相关产品推荐
相关产品推荐

