numpy使用permutation索引取值时shape异常问题咨询
NumPy重排数组时维度顺序变化问题解答
问题结论
这是NumPy的预期行为,由高级索引与基本索引混合使用时的维度排序规则导致,你使用的1.19.x属于较旧的正式版本,对这类混合索引的维度处理遵循旧版规则。
原因说明
你用到的索引a[:,shuffle_order,:,0]属于「基本索引(冒号切片、固定整数)+ 高级索引(数组shuffle_order)」的混合索引场景:
- 原数组
a的维度顺序为:轴0(size=12)、轴1(size=5)、轴2(size=2)、轴3(size=2)、轴4(size=3) - 你给出的索引依次作用于前4个轴:轴0用冒号全取、轴1用
shuffle_order数组索引、轴2用冒号全取、轴3用整数0取固定位置(该维度直接消失)、轴4默认全取 - 旧版NumPy(1.19.x及更早)的混合索引规则:当高级索引和非连续的基本索引组合使用时,会将高级索引对应的维度挪到结果数组的最前面,剩余基本索引对应的维度按原有顺序排在后面,因此就出现了你看到的
(5, 12, 2, 3)的结果。1.20及之后的新版NumPy优化了单高级索引场景的处理,你当前的写法在新版中会直接返回你预期的(12,5,2,3)结果。
修复方案
如果要得到你预期的shape,推荐两种稳妥写法,不受版本规则影响:
- 用
np.take显式指定重排的轴,维度顺序完全和原数组保持一致:
b = np.take(a, shuffle_order, axis=1)[:, :, :, 0]
- 索引后手动调整维度顺序:
b = a[:,shuffle_order,:,0].transpose(1, 0, 2, 3)
内容的提问来源于stack exchange,提问作者user3235916
相关产品推荐
相关产品推荐

