如何在PyTorch中无需循环从1D张量按起止索引提取切片
如何在PyTorch中从1D张量批量生成等长切片的2D张量
给定1D数据张量和批量切片的起始索引(切片长度固定):
import torch data = torch.arange(10) starts = torch.tensor([0, 3, 4, 1]) # 切片长度固定为2,ends = starts + 2
目标是不通过循环,直接生成如下2D张量:
tensor([[0, 1], [3, 4], [4, 5], [1, 2]])
问题原因
直接使用data[starts:ends]会报错,因为PyTorch的切片语法仅支持单个整数或单元素张量作为切片的起始/结束位置,不支持批量的起始/结束索引张量。
解决方案:利用广播生成批量索引
因为所有切片长度固定,我们可以通过广播机制生成所有切片的索引矩阵,再直接索引原张量:
slice_length = 2 # 生成每个切片内的偏移量:[0, 1],扩展维度以支持广播 offsets = torch.arange(slice_length).unsqueeze(0) # 将starts扩展维度后与偏移量广播,得到(4,2)的索引矩阵 indices = starts.unsqueeze(1) + offsets # 索引原张量得到结果 dataSlices = data[indices]
运行结果:
>>> dataSlices tensor([[0, 1], [3, 4], [4, 5], [1, 2]])
另一种方法:使用torch.gather
如果需要更灵活的索引场景,也可以用torch.gather实现:
indices = starts.unsqueeze(1) + torch.arange(slice_length) dataSlices = torch.gather(data.unsqueeze(0).repeat(len(starts), 1), 1, indices)
不过这种方法需要先重复原张量,效率略低于第一种广播索引的方式。
内容的提问来源于stack exchange,提问作者EternalTrail
相关产品推荐
相关产品推荐

