PyTorch torch.transpose Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.transposeIs a function in PyTorch used to swap two dimensions of a tensor. It returns a transposed view of the input tensor.

This is a commonly used operation in deep learning to change the shape of data to meet different computational needs.

Function Definition

torch.transpose(input, dim0, dim1)

Parameters:

  • input(Tensor): The input tensor.
  • dim0(int): The first dimension to swap.
  • dim1(int): The second dimension to swap.

Return Value:

  • torch.Tensor: Returns the transposed view of the tensor.

Usage Examples

Example 1: 2D Matrix Transpose

Example

import torch

# Create a 3x4 matrix
x = torch.randn(3, 4)

# Transpose
y = torch.transpose(x, 0, 1)

print("Original shape:", x.shape)
print("Shape after transpose:", y.shape)
print("Original:")
print(x)
print("After transpose:")
print(y)

The output is:

原始形状: torch.Size([3, 4])
转置后形状: torch.Size([4, 3])
原始:
tensor([[ 0.3364, -0.7844,  0.9760,  0.4381],
        [ 0.7865, -1.2775,  0.5767, -0.5268],
        [-0.6399, -0.6743, -0.2972, -0.4781]])
转置后:
tensor([[ 0.3364,  0.7865, -0.6399],
        [-0.7844, -1.2775, -0.6743],
        [ 0.9760,  0.5767, -0.2972],
        [ 0.4381, -0.5268, -0.4781]])

Example 2: Multi-dimensional Tensor Transpose

Example

import torch

# Create a 3D tensor
x = torch.randn(2, 3, 4)

# Swap dim=1 and dim=2
y = torch.transpose(x, 1, 2)

print("Original shape:", x.shape)
print("Shape after transpose:", y.shape)

The output is:

原始形状: torch.Size([2, 3, 4])
转置后形状: torch.Size([2, 4, 3])

Notes

  • torch.transposeIt returns a view, not a copy.
  • For 2D tensors, you can also use thetensor.t()method.

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions