NumPy中使用ndarray索引数组时np.take结果不符合预期如何解决
numpy.take沿指定轴索引多维数组的用法说明 问题原因
你调用np.take时没有指定axis参数,函数默认会将原数组展平为一维数组后再执行索引操作,因此返回结果和预期不符。
你当前写的np.take(A, idx)逻辑等价于A.ravel()[idx]:将形状为(5,5,3)的A展平为长度75的一维数组,再用idx中的值作为一维索引取值。比如idx[1,1]的值是1,最终取到的是展平数组中索引为1的元素1,而非你期望的第三维度索引为1的位置值。
正确写法
numpy.take沿指定轴索引时,必须通过axis参数明确指定要操作的轴。你需要沿第三维度(NumPy轴序号从0开始计数,第三维度对应axis=2)索引,代码修改如下:
import numpy as np # 构造目标矩阵 A = np.arange(75).reshape((5,5,3)) # 构造索引数组 idx = np.array([[1, 0, 0, 1, 1], [1, 1, 0, 1, 1], [1, 0, 1, 0, 1], [1, 1, 0, 0, 0], [1, 1, 1, 1, 0]]) # 沿第三维度(axis=2)用idx取值 Asub = np.take(A, idx, axis=2) # 结果验证 print(f'A在[1,1,1]位置的值是 {A[1,1,1]}') print(f'idx在[1,1]位置存储的索引值是 {idx[1,1]}') print(f'Asub在[1,1]位置的值是 {Asub[1,1]}')
运行后输出符合预期:
A在[1,1,1]位置的值是 19 idx在[1,1]位置存储的索引值是 1 Asub在[1,1]位置的值是 19
补充说明
- 当
np.take的axis参数为默认值None时,永远是对展平后的一维数组做索引,使用时如果需要操作多维数组的特定轴,必须显式传入axis参数。 - 该场景也可以通过NumPy高级索引实现,效果和指定
axis=2的take完全一致:
i, j = np.ogrid[:A.shape[0], :A.shape[1]] Asub = A[i, j, idx]
内容的提问来源于stack exchange,提问作者brechmos
相关产品推荐
相关产品推荐

