You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 21:54:59