如何高效从三维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
相关产品推荐
相关产品推荐

