NumPy用数组argmax索引另一数组取值报错解决方法
问题描述
需要先计算4维ndarray在第0维的argmax结果,再用该索引值从另一个首维度长度一致、其余维度不同的4维ndarray中提取对应位置的值,目标输出形状为(1,12,100,100)。
原实现代码如下:
import numpy as np ndvi_array = np.random.randint(0, 255, size=(4, 1, 100, 100)) image_array = np.random.randint(0, 255, size=(4, 12, 100, 100)) ndvi_argmax = ndvi_array.argmax(0) print(f"NDVI argmax shape: {ndvi_argmax.shape}") zipped = tuple(zip(range(len(ndvi_argmax)), ndvi_argmax)) result = image_array[zipped] print(f"Result shape: {result.shape}")
运行代码抛出如下错误:
only integers, slices (`:`), ellipsis (`...`), numpy.newaxis (`None`) and integer or boolean arrays are valid indices
报错原因
numpy多维高级索引要求传入的索引序列是若干个形状可广播的整数/布尔数组,每个数组对应一个维度的索引位置。原代码中手动zip得到的是嵌套元组结构,不符合numpy索引的格式要求,无法被识别为合法的多维索引。
实际上ndvi_argmax本身已经是形状为(1,100,100)的整数数组,存储了每个(通道, 高, 宽)位置上ndvi最大值对应的第0维下标,只需要配合其余维度的广播索引即可完成取值,不需要手动拼接元组。
正确实现
方法1:使用np.take_along_axis(最简洁)
直接沿第0维按索引取值,只需要给ndvi_argmax插入一个长度为1的维度匹配image_array的通道维广播规则即可:
import numpy as np ndvi_array = np.random.randint(0, 255, size=(4, 1, 100, 100)) image_array = np.random.randint(0, 255, size=(4, 12, 100, 100)) ndvi_argmax = ndvi_array.argmax(0) # 在通道维插入新轴,将索引形状从(1,100,100)转为(1,1,100,100),适配广播规则 result = np.take_along_axis(image_array, ndvi_argmax[:, np.newaxis, :, :], axis=0) print(result.shape) # 输出 (1, 12, 100, 100),符合预期
方法2:手动构造网格索引
使用np.ogrid生成其余维度的开放索引网格,和ndvi_argmax组合成合法的高级索引序列:
import numpy as np ndvi_array = np.random.randint(0, 255, size=(4, 1, 100, 100)) image_array = np.random.randint(0, 255, size=(4, 12, 100, 100)) ndvi_argmax = ndvi_array.argmax(0) # 生成第1、2、3维的开放网格,形状可广播匹配 _, c, h, w = np.ogrid[:4, :12, :100, :100] result = image_array[ndvi_argmax, c, h, w] print(result.shape) # 输出 (1, 12, 100, 100)
内容的提问来源于stack exchange,提问作者Vitaly Olegovitch
相关产品推荐
相关产品推荐

