Numpy高维数组用列表与numpy数组高级索引结果差异原因
NumPy 嵌套列表与NumPy数组索引行为差异原因
两者返回结果不同,核心是NumPy对两类索引输入的解析逻辑完全不同:
- 传入嵌套Python列表做索引时,NumPy会做兼容解析,将嵌套结构识别为对应不同轴的坐标索引序列。你定义的
indices = [[0, 1], [1, 0], [1, 1]]包含3个子列表,和目标数组a的3个维度一一对应,分别匹配第0、1、2轴的取数下标,等价于执行a[[0,1], [1,0], [1,1]],最终取到a[0,1,1]=2、a[1,0,1]=1,正好是预期结果。 - 传入NumPy数组做索引时,数组本身是独立的合法索引对象,NumPy默认不会自动拆分它的维度去匹配多个轴,只会把它当作第0轴的高级索引源处理。你的
indices_arr是shape为(3,2)的二维数组,所有元素值都是0或1(正好是a第0轴仅有的两个合法下标),NumPy会按索引数组的每个元素值取第0轴对应的二维子块,再按索引数组的形状拼接,最终得到shape为(3,2,2,2)的四维结果,和预期完全不符。
正确实现方式
如果要用NumPy数组实现逐点坐标取数,需要手动把索引数组拆为对应各轴的形式,最简便的写法是将索引数组转置后转为元组传入:
import numpy as np a = np.array([ [ [1, 3], [0, 2] ], [ [2, 1], [4, 2] ] ], dtype=np.int32) indices_arr = np.array([[0, 1], [1, 0], [1, 1]], dtype=np.int32) # 转置后每一列对应一个轴的坐标,转元组后会被识别为分轴索引 print(a[tuple(indices_arr.T)]) # 输出: array([2, 1], dtype=int32)
核心规则:NumPy仅在索引位置传入元组时,才会将元组的每个元素匹配到对应维度的索引;传入单个Python列表时会触发兼容逻辑,按分轴坐标解析;传入单个NumPy数组时,会默认作为第0轴的高级索引处理。
内容的提问来源于stack exchange,提问作者Max
相关产品推荐
相关产品推荐

