Numpy indexing与broadcast应用:高维张量指定位置元素批量选取方法
解法
你可以通过构造三个可广播对齐的索引数组,一次性完成取值操作,核心写法如下:
import numpy as np # 示例张量a构造 a = np.array([ [[-1.054, 0.068, -0.572, 1.535, 1.746], [-0.115, 0.356, 0.222, -0.391, 0.367], [-0.53 , -0.856, 0.58 , 1.099, 0.605]], [[ 0.31 , 0.037, -0.85 , -0.054, -0.75 ], [-0.097, -1.707, -0.702, 0.658, 0.548], [ 1.727, -0.326, -1.525, -0.656, 0.349]] ]) # 一次性索引写法 res = a[ [[0], [1]], range(3), [[0,2,4], [1,3,2]] ] print(res)
运行输出结果和分开计算的结果完全一致:
[[-1.054 0.222 0.605] [ 0.037 0.658 -1.525]]
原理说明
numpy的高级索引要求各维度的索引数组形状可广播对齐:
- 第一维度的索引
[[0],[1]]形状为(2, 1),对应要取的第0、第1个3×5子张量 - 第二维度的索引
range(3)(即[0,1,2])形状为(3,),对应每个子张量的3行 - 第三维度的索引
[[0,2,4], [1,3,2]]形状为(2, 3),对应每行要取的列位置
三者广播后形状都对齐为(2, 3),每个位置的三组索引刚好对应目标元素位置。
如果需要适配更高维度的批量场景,可以用更通用的写法生成索引:
batch_idx = np.arange(a.shape[0])[:, np.newaxis] # 自动生成形状为(batch_size, 1)的批次索引 row_idx = np.arange(a.shape[1]) col_idx = np.array([[0,2,4], [1,3,2]]) res = a[batch_idx, row_idx, col_idx]
内容的提问来源于stack exchange,提问作者Sam-gege
相关产品推荐
相关产品推荐

