如何在NumPy中实现基于可变长度元组的动态多维索引?
NumPy 多维数组动态索引问题
给定以下简化示例:
import numpy as np z = np.arange(125).reshape(5,5,5) i = np.where( np.mod(z[0], 3) == 0 )
此时i是包含2个数组的元组,可用于单个切片的索引:
z[0][i] # array([ 0, 3, 6, 9, 12, 15, 18, 21, 24]) z[1][i] # array([25, 28, 31, 34, 37, 40, 43, 46, 49])
需要对z的所有切片执行该索引操作:
- 尝试
z[:, i]得到的形状为(5, 2, 9, 5),不符合预期的(5, 9); - 已知
z[:, i[0], i[1]]可得到正确结果,但实际场景中i包含的数组数量不确定,无法硬编码索引; - 尝试
z[:, *i]会触发SyntaxError。
问题:能否实现类似z[:, *i]的效果,即当i包含n个数组时,等价于z[:, i[0], i[1], ..., i[n-1]]?
解决方案
可以通过构造动态索引元组实现需求,核心思路是将第一个维度的切片(:)与i中的元素合并成一个新的元组,直接用于索引。
方法1:拼接索引元组
利用元组拼接特性,把slice(None)(等价于:)和i合并:
idx = (slice(None),) + i result = z[idx] print(result.shape) # (5, 9) print(result) # array([[ 0, 3, 6, 9, 12, 15, 18, 21, 24], # [25, 28, 31, 34, 37, 40, 43, 46, 49], # [50, 53, 56, 59, 62, 65, 68, 71, 74], # [75, 78, 81, 84, 87, 90, 93, 96, 99], # [100, 103, 106, 109, 112, 115, 118, 121, 124]])
方法2:动态生成索引列表再转元组
如果需要更灵活的构造方式,也可以先把切片转为列表元素,再和i的列表形式合并,最后转成元组:
idx = tuple([slice(None)] + list(i)) result = z[idx]
原理说明
z[:, i]不符合预期的原因:NumPy会把元组i当作第二个维度的索引对象,而非将元组内的元素分配到后续维度,因此会引入额外维度;- 拼接后的元组
idx会被NumPy解析为(slice(None), i[0], i[1], ..., i[n-1]),完全等价于手动硬编码的索引方式,且不受i中元素数量的限制。
内容的提问来源于stack exchange,提问作者evanb
相关产品推荐
相关产品推荐

