如何简化PyTorch中3D张量的切片操作
PyTorch批量切片3D张量(移除for循环)
需求说明
对形状为(batch, max_len, hidden_dim)的3D张量src_tensor,在第二维度上按(batch,)形状的索引向量indices批量切片:每个样本取第二维度中indices[i]和indices[i]+1位置的两个元素,替代原有的for循环实现。
原实现代码
import torch nums = 30 l = [i for i in range(nums)] src_tensor = torch.Tensor(l).reshape((3,5,2)) indices = [1,2,3] slice_tensor = torch.zeros((3,2,2)) for i in range(3): p1,p2 = indices[i],indices[i]+1 slice_tensor[i,:,:]=src_tensor[i,[p1,p2],:] print(src_tensor) print(indices) print(slice_tensor)
优化后代码(移除for循环)
import torch nums = 30 l = [i for i in range(nums)] src_tensor = torch.Tensor(l).reshape((3,5,2)) indices = torch.tensor([1,2,3]) # 转换为张量便于运算 # 生成每个样本对应的第二维度索引:每个起始位置+0、+1 slice_indices = indices.unsqueeze(1) + torch.arange(2) # 形状变为(3,2) # 利用高级索引批量取数 batch_size = src_tensor.shape[0] slice_tensor = src_tensor[torch.arange(batch_size), slice_indices, :] # 输出验证 print(src_tensor) print(indices) print(slice_tensor)
原理说明
- 索引生成:将
indices转为张量后,通过unsqueeze(1)扩展维度为(batch,1),与torch.arange(2)(即[0,1])做加法,得到每个样本需要取的两个第二维度索引,形状为(batch,2); - 高级索引:用
torch.arange(batch_size)匹配每个batch样本,slice_indices对应每个样本的第二维度位置,直接从src_tensor中批量提取目标元素,无需预先初始化零张量,完全通过向量化操作替代循环,效率更高。
输出结果与原代码完全一致。
内容的提问来源于stack exchange,提问作者MasterLu
相关产品推荐
相关产品推荐

