PyTorch torch.movedim Function
Pytorch torch Reference Manual
torch.movedimIt is a function in PyTorch for moving tensor dimensions. It moves the specified dimensions to new positions and returns a new tensor view.
This function is very useful when adjusting the shape of a tensor for specific operations, such as when processing image data or preparing neural network inputs.
Function Definition
torch.movedim(input, source, destination)
Parameters:
input(Tensor): The input tensor.source(int or tuple of int): The original dimension indices to move. Can be an integer or a tuple of dimension indices.destination(int or tuple of int): The target position indices. Can be an integer or a tuple of dimension indices, with the same length as source.
Return Value:
torch.Tensor: Returns the tensor view after the dimensions are moved.
Usage Examples
Example
import torch
# Create a tensor with shape (batch, channel, height, width)
x = torch.randn(32, 3, 224, 224)
# Move the channel dimension to the last dimension
# Move from position 1 to position 3
y = torch.movedim(x, 1, 3)
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
# Create a tensor with shape (batch, channel, height, width)
x = torch.randn(32, 3, 224, 224)
# Move the channel dimension to the last dimension
# Move from position 1 to position 3
y = torch.movedim(x, 1, 3)
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
The output result is:
原始形状: torch.Size([32, 3, 224, 224]) 移动后形状: torch.Size([32, 224, 224, 3])
Example
import torch
# Create a tensor with shape (D0, D1, D2, D3)
x = torch.randn(2, 3, 4, 5)
# Move multiple dimensions at once
# Move dimensions 0 and 1 to positions 2 and 3
y = torch.movedim(x, source=(0, 1), destination=(2, 3))
print("Original shape:", x.shape)
print("Shape after moving:", y.shape)
# Create a tensor with shape (D0, D1, D2, D3)
x = torch.randn(2, 3, 4, 5)
# Move multiple dimensions at once
# Move dimensions 0 and 1 to positions 2 and 3
y = torch.movedim(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])
Example
import torch
# Convert an image tensor from (N, C, H, W) to (N, H, W, C)
# This is useful when sending data to APIs that require the channel in the last position
images = torch.randn(16, 3, 64, 64)
images_permuted = torch.movedim(images, 1, 3)
print("Original shape (N, C, H, W):", images.shape)
print("Converted shape (N, H, W, C):", images_permuted.shape)
# Convert an image tensor from (N, C, H, W) to (N, H, W, C)
# This is useful when sending data to APIs that require the channel in the last position
images = torch.randn(16, 3, 64, 64)
images_permuted = torch.movedim(images, 1, 3)
print("Original shape (N, C, H, W):", images.shape)
print("Converted shape (N, H, W, C):", images_permuted.shape)
Note:torch.movedimWhat is returned is a view, not a copy, so the operation is efficient. There is also a similar functiontorch.moveaxis, whichtorch.movedimhas the same functionality, but the parameter names are different.
Other Extensions