PyTorch torch.reshape Function


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

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)

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)

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())

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 usetensor.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.

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions