PyTorch torch.movedim Function


Pytorch torch 参考手册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)

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)

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)

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.


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions