PyTorch torch.moveaxis Function
Pytorch torch Reference Manual
torch.moveaxisis a function in PyTorch used to move tensor axes. It moves the specified axes to new positions and returns a new tensor view.
This function is the same astorch.movedimIt has the same functionality, except the parameter names differ (moveaxis uses "axis" while movedim uses "dim").
Function Definition
torch.moveaxis(input, source, destination)
Parameters:
input(Tensor): The input tensor.source(int or tuple of int): The original axis indices to move. Can be an integer or a tuple of axis indices.destination(int or tuple of int): The target position indices. Can be an integer or a tuple of axis indices, and its length must be the same as source.
Return Value:
torch.Tensor: Returns the tensor view after moving the axes.
Usage Examples
Example
import torch
# Create a tensor with shape (batch, seq_len, feature)
x = torch.randn(32, 10, 128)
# Move the last dimension to the first position
y = torch.moveaxis(x, -1, 0)
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
# Create a tensor with shape (batch, seq_len, feature)
x = torch.randn(32, 10, 128)
# Move the last dimension to the first position
y = torch.moveaxis(x, -1, 0)
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
The output result is:
原始形状: torch.Size([32, 10, 128]) 移动后形状: torch.Size([128, 10, 32])
Example
import torch
# Create a four-dimensional tensor (time, channel, height, width)
x = torch.randn(10, 3, 32, 32)
# Move the time dimension to the end
y = torch.moveaxis(x, 0, -1)
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
# Create a four-dimensional tensor (time, channel, height, width)
x = torch.randn(10, 3, 32, 32)
# Move the time dimension to the end
y = torch.moveaxis(x, 0, -1)
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
The output result is:
原始形状: torch.Size([10, 3, 32, 32]) 移动后形状: torch.Size([3, 32, 32, 10])
Example
import torch
# Move multiple axes at once
x = torch.randn(2, 3, 4, 5)
# Move axes 0 and 1 to positions 2 and 3 together
y = torch.moveaxis(x, source=(0, 1), destination=(2, 3))
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
# Move multiple axes at once
x = torch.randn(2, 3, 4, 5)
# Move axes 0 and 1 to positions 2 and 3 together
y = torch.moveaxis(x, source=(0, 1), destination=(2, 3))
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
The output result is:
原始形状: torch.Size([2, 3, 4, 5]) 移动后形状: torch.Size([4, 5, 2, 3])
Note:torch.moveaxisIt returns a view, not a copy, so the operation is efficient. This function is often used to adjust data shapes to suit different deep learning frameworks or APIs.
Other Extensions