PyTorch torch.reshape Function
Pytorch torch Reference Manual
torch.reshapeIt is a function in PyTorch used to change the shape of a tensor. It returns a new tensor with the same number of elements as the original tensor, but with a different shape.
This is a very common operation in deep learning, used to adjust the data shape to meet the input requirements of different layers.
Function Definition
torch.reshape(input, shape)
Parameters:
input(Tensor): The input tensor.shape(tuple or int): The target shape. The number of elements in the shape must equal the number of elements in the input tensor. You can use-1to let PyTorch automatically infer the dimension.
Return Value:
torch.Tensor: Returns a tensor view with the changed shape.
Usage Examples
Example 1: Reshape a 1D tensor to 2D
Example
import torch
# Create a 1D tensor
x = torch.arange(12)
print("Original shape:", x.shape)
print(x)
# Change to a 3x4 2D tensor
y = torch.reshape(x, (3, 4))
print("New shape:", y.shape)
print(y)
# Create a 1D tensor
x = torch.arange(12)
print("Original shape:", x.shape)
print(x)
# Change to a 3x4 2D tensor
y = torch.reshape(x, (3, 4))
print("New shape:", y.shape)
print(y)
The output is:
原始形状: torch.Size([12])
tensor([ 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11])
新形状: torch.Size([3, 4])
tensor([[ 0, 1, 2, 3],
[ 4, 5, 6, 7],
[ 8, 9, 10, 11]])
Example 2: Using -1 to automatically infer dimensions
Example
import torch
# Create a 3D tensor
x = torch.randn(2, 3, 4)
print("Original shape:", x.shape)
# Use -1 to automatically infer the last dimension
y = torch.reshape(x, (2, -1))
print("New shape:", y.shape)
# Create a 3D tensor
x = torch.randn(2, 3, 4)
print("Original shape:", x.shape)
# Use -1 to automatically infer the last dimension
y = torch.reshape(x, (2, -1))
print("New shape:", y.shape)
The output is:
原始形状: torch.Size([2, 3, 4]) 新形状: torch.Size([2, 12])
Example 3: Flatten to 2D
Example
import torch
# Create a 4D tensor (typical batched image data)
x = torch.randn(32, 3, 224, 224) # Batch size 32, channels 3, height and width 224
print("Original shape:", x.shape)
# Flatten to (batch size, other)
y = torch.reshape(x, (32, -1))
print("Shape after flattening:", y.shape)
# Create a 4D tensor (typical batched image data)
x = torch.randn(32, 3, 224, 224) # Batch size 32, channels 3, height and width 224
print("Original shape:", x.shape)
# Flatten to (batch size, other)
y = torch.reshape(x, (32, -1))
print("Shape after flattening:", y.shape)
The output is:
原始形状: torch.Size([32, 3, 224, 224]) 展平后形状: torch.Size([32, 150528])
This is often used before feeding image data into fully connected layers.
Example 4: Difference between reshape and view
Example
import torch
# Create a contiguous tensor
x = torch.arange(12).reshape(3, 4)
# reshape may return a view or a copy
y = torch.reshape(x, (4, 3))
print("y is a view of x:", y.is_contiguous() or y.data_ptr() == x.data_ptr())
# Create a contiguous tensor
x = torch.arange(12).reshape(3, 4)
# reshape may return a view or a copy
y = torch.reshape(x, (4, 3))
print("y is a view of x:", y.is_contiguous() or y.data_ptr() == x.data_ptr())
torch.reshapecan handle non-contiguous tensors, whiletensor.view()requires that the tensor must be contiguous.
Notes
torch.reshapeThe returned tensor may be a view of the original data or a copy, depending on the memory layout.- If you need to guarantee a returned view, you can use
tensor.view(), but make sure the tensor is contiguous. - The total number of elements in the shape must be the same as the number of elements in the original tensor.
Other Extensions