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

如何高效从三维PyTorch Tensor按指定索引提取子张量?

高效提取三维Tensor/数组指定索引的子数组(PyTorch/NumPy方案)

需求描述

我有一个形状为(a, b, c)的三维数组/Tensor,还有一个长度为a的索引列表B,每个索引取值范围是[0, b)。需要得到形状为(a, c)的数组,目前用列表推导实现:

z = torch.stack([t_[b, :] for t_, b in zip(tensor, B)])

这段代码用于神经网络前向传播,希望避免列表推导,寻求更高效的PyTorch或NumPy实现方式。

背景:处理不同长度时间窗口的时序数据,用PyTorch的pack_padded_sequence及逆操作做掩码处理,需要获取掩码开始前LSTM的输出(后续网络输出失效)。

示例

# 输入Tensor,shape: (4, 3, 2)
tensor = [[[ 0,  1],
           [ 2,  3],
           [ 4,  5]],
          [[ 6,  7],
           [ 8,  9],
           [10, 11]],
          [[12, 13],
           [14, 15],
           [16, 17]],
          [[18, 19],
           [20, 21],
           [22, 23]]]
B = [0, 1, 2, 2]

# 期望输出,shape: (4, 2)
output = [[ 0,  1],
          [ 8,  9],
          [16, 17],
          [22, 23]]

解决方案

PyTorch实现

使用高级索引直接完成向量化提取,无需循环,效率更高且能保留计算图(适配前向传播需求):

import torch

tensor = torch.tensor(tensor)
B = torch.tensor(B)

# 核心代码
z = tensor[torch.arange(tensor.shape[0]), B, :]

逻辑说明:

  • torch.arange(tensor.shape[0])生成第一个维度的索引序列[0,1,2,3],对应每个样本
  • B是每个样本在第二个维度的目标索引
  • 两者组合后,会为每个样本选取tensor[i, B[i], :],最终拼接成(a,c)的Tensor

NumPy实现

逻辑与PyTorch一致,用NumPy的高级索引完成:

import numpy as np

tensor = np.array(tensor)
B = np.array(B)

# 核心代码
z = tensor[np.arange(tensor.shape[0]), B, :]

优势对比

相比列表推导,高级索引的优势:

  • 完全向量化操作,避免Python循环,处理大张量时速度提升明显
  • PyTorch版本能被自动微分机制追踪,不会破坏计算图,适配神经网络训练场景

内容的提问来源于stack exchange,提问作者mhenning

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 14:12:27