PyTorch torch.transpose Function
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)
# 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)
# 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 the
tensor.t()method.
Other Extensions