Numpy多维索引疑问:元组索引与数组索引结果为何不同?
问题分析与解决
你遇到的核心问题是numpy索引的广播规则以及np.diag_indices的参数细节,具体拆解如下:
手动使用
m[:, np.array([0,1,2,3]), np.array([0,1,2,3])]时,这是对应位置的花式索引:numpy会为每个矩阵选取(0,0)、(1,1)、(2,2)、(3,3)四个对角线元素,最终得到形状为(10,4)的数组,符合预期。若
m[:, tuple(din)]返回原数组(形状(10,4,4)),大概率是调用np.diag_indices时漏掉了ndim=2参数:- 默认情况下
np.diag_indices(n)的ndim=1,返回单元素元组(array([0,1,2,3]),),仅包含行索引。 - 此时
m[:, tuple(din)]等价于m[:, [0,1,2,3], :],即选取每个矩阵的所有行和列,结果形状与原数组一致,看起来像是返回了原数组。
- 默认情况下
当你正确指定ndim=2时,np.diag_indices(m.shape[1], ndim=2)会返回包含行、列索引的二元组(array([0,1,2,3]), array([0,1,2,3])),此时m[:, tuple(din)]和手动索引的结果完全一致,都会返回每个矩阵的对角线元素数组。
验证代码:
import numpy as np m = np.random.normal(0, 0.2, (10, 4, 4)) # 正确获取二维对角线索引 din = np.diag_indices(m.shape[1], ndim=2) print(m[:, tuple(din)].shape) # 输出 (10,4),符合预期 # 错误调用(无ndim=2) din_wrong = np.diag_indices(m.shape[1]) print(m[:, tuple(din_wrong)].shape) # 输出 (10,4,4),与原数组形状一致
若代码确实写了ndim=2但结果异常,检查两点:
- 打印
din确认内容:应为包含两个长度为4的数组的元组。 - 升级numpy到1.20+版本,避免旧版本的索引行为差异。
内容的提问来源于stack exchange,提问作者prodx
相关产品推荐
相关产品推荐

