PyTorch torch.roll Function
PyTorch torch Reference Manual
torch.rollThis is a function in PyTorch used to roll tensor elements.
Function Definition
torch.roll(input, shifts, dims)
Usage Example
Example
import torch
x = torch.arange(8)
# Roll forward by 2 positions
print(torch.roll(x, 2))
# Roll backward by 2 positions
print(torch.roll(x, -2))
x = torch.arange(8)
# Roll forward by 2 positions
print(torch.roll(x, 2))
# Roll backward by 2 positions
print(torch.roll(x, -2))
The output result is:
tensor([6, 7, 0, 1, 2, 3, 4, 5]) tensor([2, 3, 4, 5, 6, 7, 0, 1])
Other Extensions