如何实现一维PyTorch Tensor元素向上滚动?
PyTorch一维Tensor反向滚动的实现方案
需求明确
现有一维PyTorch Tensor:
import torch tensor = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
需要实现反向滚动(元素整体向右偏移一位),基础目标结果为:[2.0, 3.0, 4.0, 5.0, 6.0, 1.0];实际场景中需将最后一个元素置零,最终结果为:[2.0, 3.0, 4.0, 5.0, 6.0, 0.0]。
注:torch.inverse是矩阵求逆方法,与元素反转无关;若要完全反转Tensor,应使用torch.flip(tensor, dims=[0]),但这和需求的反向滚动逻辑不同。
实现方案
方案一:利用torch.roll负偏移量+索引置零
torch.roll的shifts参数为负数时,会实现反向滚动(向右偏移),之后直接修改最后一个元素即可。
tensor = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) # 反向滚动一位 rolled_tensor = torch.roll(tensor, shifts=-1) # 最后一位置零 rolled_tensor[-1] = 0.0 print(rolled_tensor) # 输出: tensor([2., 3., 4., 5., 6., 0.])
方案二:切片拼接直接生成结果
通过Tensor切片提取后半部分,再拼接值为0的单元素Tensor,一步到位得到最终结果。
tensor = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) result = torch.cat([tensor[1:], torch.tensor([0.0])]) print(result) # 输出: tensor([2., 3., 4., 5., 6., 0.])
方案三:torch.roll结合scatter_批量置零(适合多元素修改场景)
如果需要批量修改指定位置的值,可使用scatter_方法替代直接索引:
tensor = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) result = torch.roll(tensor, shifts=-1).scatter_(0, torch.tensor([5]), 0.0) print(result) # 输出: tensor([2., 3., 4., 5., 6., 0.])
内容的提问来源于stack exchange,提问作者FlumeRS
相关产品推荐
相关产品推荐

