PyTorch张量索引的高效实现方法求助
PyTorch高效张量索引实现
问题描述
现有形状为(2,2,2,4)的张量X和形状为(4,)的索引向量Y,具体定义如下:
import torch X=torch.tensor([[[[8, 2, 8, 5], [3, 7, 4, 0]], [[4, 5, 7, 4], [8, 3, 9, 5]]], [[[5, 2, 9, 3], [6, 4, 5, 4]], [[7, 3, 3, 7], [6, 3, 8, 9]]]]) Y=torch.tensor([1,1,0,1])
需要用Y对X进行索引,得到目标张量。手动实现方式为:
torch.stack([X[1][0][0], X[1][0][1], X[0][1][0], X[1][1][1]])
目标结果为:
tensor([[5, 2, 9, 3], [6, 4, 5, 4], [4, 5, 7, 4], [6, 3, 8, 9]])
目前通过for循环实现效率极低,需要PyTorch中的高效实现方案。
高效实现方案
直接使用PyTorch的高级索引完成向量化操作,完全替代循环,效率拉满。
步骤分析
先拆解手动索引的规律:
你要取的4个元素对应:
X[Y[0], 0, 0]→Y[0]=1→X[1,0,0]X[Y[1], 0, 1]→Y[1]=1→X[1,0,1]X[Y[2], 1, 0]→Y[2]=0→X[0,1,0]X[Y[3], 1, 1]→Y[3]=1→X[1,1,1]
其中后两个维度的索引是(0,0), (0,1), (1,0), (1,1),也就是2×2网格的展平序列。
代码实现
针对当前固定维度的写法
直接生成后两个维度的索引,配合Y批量取值:
# 定义后两个维度的展平索引 dim1 = torch.tensor([0, 0, 1, 1]) dim2 = torch.tensor([0, 1, 0, 1]) # 高级索引批量获取结果 result = X[Y, dim1, dim2] # 输出验证 print(result)
运行结果与手动实现完全一致:
tensor([[5, 2, 9, 3], [6, 4, 5, 4], [4, 5, 7, 4], [6, 3, 8, 9]])
通用扩展写法
如果后续后两个维度的大小变化(比如不是2×2),可以用torch.meshgrid自动生成索引,无需手动硬编码:
# 假设X的第2、3维度是(M,N),这里M=N=2 M, N = X.shape[1], X.shape[2] # 生成网格索引并展平 dim1, dim2 = torch.meshgrid(torch.arange(M), torch.arange(N), indexing='ij') dim1 = dim1.flatten() dim2 = dim2.flatten() # 批量索引 result = X[Y, dim1, dim2]
这种写法适配任意大小的中间维度,复用性更强。
内容的提问来源于stack exchange,提问作者Frank Furt
相关产品推荐
相关产品推荐

