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

如何简化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)

原理说明

  1. 索引生成:将indices转为张量后,通过unsqueeze(1)扩展维度为(batch,1),与torch.arange(2)(即[0,1])做加法,得到每个样本需要取的两个第二维度索引,形状为(batch,2);
  2. 高级索引:用torch.arange(batch_size)匹配每个batch样本,slice_indices对应每个样本的第二维度位置,直接从src_tensor中批量提取目标元素,无需预先初始化零张量,完全通过向量化操作替代循环,效率更高。

输出结果与原代码完全一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 19:03:12