如何通过一次操作从PyTorch多维张量中按指定索引提取元素?
解决方案
可以通过高级索引直接对原张量执行一次操作得到目标结果,核心是构造匹配需求的索引组合:
方式一:直接拼接成一维序列(对应分步操作后的拼接结果)
import torch arr = torch.randint(0, 9, (100, 50, 3)) # 构造第一维度索引:重复每个目标样本的索引,次数等于要提取的元素数量 idx_dim0 = torch.tensor([5]*6 + [55]*6) # 构造第二维度索引:分别对应两个样本的提取范围(左闭右开,5-10对应5:11,10-15对应10:16) idx_dim1 = torch.tensor([5,6,7,8,9,10] + [10,11,12,13,14,15]) # 一次索引得到最终结果 final_result = arr[idx_dim0, idx_dim1]
最终结果形状为 torch.Size([12, 3]),和分步操作(先取partial_arr、分别切片后拼接)的输出完全一致。
方式二:保留(2,6,3)的样本维度结构
如果需要保留两个样本的独立维度,可以通过掩码索引实现:
import torch arr = torch.randint(0, 9, (100, 50, 3)) indices = torch.tensor([5, 55]) # 定义每个样本的切片起止位置 starts = torch.tensor([5, 10]) ends = torch.tensor([11, 16]) # 生成第二维度的全范围索引,并匹配每个样本的提取范围 idx_dim1 = torch.arange(50).unsqueeze(0).repeat(2, 1) mask = (idx_dim1 >= starts.unsqueeze(1)) & (idx_dim1 < ends.unsqueeze(1)) # 一次索引并整理形状 final_result = arr[indices.unsqueeze(1).repeat(1,6), idx_dim1[mask]].reshape(2,6,3)
内容的提问来源于stack exchange,提问作者kklaw
相关产品推荐
相关产品推荐

