PyTorch torch.expand Function
Pytorch torch Reference ManualPytorch torch Reference Manual
torch.expandis a function in PyTorch used to expand tensor dimensions. It expands the tensor by duplicating the view without actually copying the data.
The expanded tensor shares the original data as a view in memory, making this a memory-efficient operation.
Function Definition
torch.expand(*sizes) torch.expand(input, *sizes)
Parameters:
input(Tensor): The input tensor.*sizes(torch.Size or int): The target size. A single dimension in the size can be -1, meaning that dimension is kept unchanged. Dimensions with a value of 1 can be expanded to larger sizes.
Return Value:
torch.Tensor: Returns the expanded view of the tensor.
Usage Examples
Example
import torch
# Create a tensor with dimension 1
x = torch.tensor([[1], [2], [3]])
print("Original tensor:")
print(x)
print("Shape:", x.shape)
# Expand to a larger size
y = x.expand(3, 4)
print("nAfter expansion:")
print(y)
print("Shape:", y.shape)
# Create a tensor with dimension 1
x = torch.tensor([[1], [2], [3]])
print("Original tensor:")
print(x)
print("Shape:", x.shape)
# Expand to a larger size
y = x.expand(3, 4)
print("nAfter expansion:")
print(y)
print("Shape:", y.shape)
The output is:
原始张量:
tensor([[1],
[2],
[3]])
形状: torch.Size([3, 1])
扩展后:
tensor([[1, 1, 1, 1],
[2, 2, 2, 2],
[3, 3, 3, 3]])
形状: torch.Size([3, 4])
Example
import torch
# Use -1 to keep the dimension unchanged
x = torch.tensor([1, 2, 3, 4])
# Expand to a 2x4 tensor
y = x.expand(2, -1)
print("Original:", x.shape)
print("After expansion:", y.shape)
print(y)
# Use -1 to keep the dimension unchanged
x = torch.tensor([1, 2, 3, 4])
# Expand to a 2x4 tensor
y = x.expand(2, -1)
print("Original:", x.shape)
print("After expansion:", y.shape)
print(y)
The output is:
原始: torch.Size([4])
扩展后: torch.Size([2, 4])
tensor([[1, 2, 3, 4],
[1, 2, 3, 4]])
Example
import torch
# Broadcasting mechanism use case
# Expand a column vector to a matrix
column = torch.randn(5, 1)
# Expand to add with a matrix
matrix = torch.randn(5, 10)
# Broadcasting: the column vector will automatically expand
result = matrix + column
print("Column vector shape:", column.shape)
print("Matrix shape:", matrix.shape)
print("Result shape:", result.shape)
# Broadcasting mechanism use case
# Expand a column vector to a matrix
column = torch.randn(5, 1)
# Expand to add with a matrix
matrix = torch.randn(5, 10)
# Broadcasting: the column vector will automatically expand
result = matrix + column
print("Column vector shape:", column.shape)
print("Matrix shape:", matrix.shape)
print("Result shape:", result.shape)
The output is:
列向量形状: torch.Size([5, 1]) 矩阵形状: torch.Size([5, 10]) 结果形状: torch.Size([5, 10])
Example
import torch
# Expand from a 2D tensor to 3D
x = torch.tensor([[1, 2, 3],
[4, 5, 6]])
# Expand the third dimension
y = x.expand(2, 2, 3)
print("Original shape:", x.shape)
print("Expanded shape:", y.shape)
print("nExpanded tensor:")
print(y)
# Verify that the data is shared
print("nModified original tensor:")
x[0, 0] = 100
print("x[0,0]:", x[0, 0])
print("y[0,0,0]:", y[0, 0, 0])
# Expand from a 2D tensor to 3D
x = torch.tensor([[1, 2, 3],
[4, 5, 6]])
# Expand the third dimension
y = x.expand(2, 2, 3)
print("Original shape:", x.shape)
print("Expanded shape:", y.shape)
print("nExpanded tensor:")
print(y)
# Verify that the data is shared
print("nModified original tensor:")
x[0, 0] = 100
print("x[0,0]:", x[0, 0])
print("y[0,0,0]:", y[0, 0, 0])
The output is:
原始形状: torch.Size([2, 3])
扩展后形状: torch.Size([2, 2, 3])
扩展后的张量:
tensor([[[1, 2, 3],
[4, 5, 6]],
[[1, 2, 3],
[4, 5, 6]]])
修改原始张量:
x[0,0]: tensor(100)
y[0,0,0]: tensor(100)
Note:torch.expandOnly dimensions with a size of 1 can be expanded to larger sizes. Dimensions larger than 1 cannot be shrunk or expanded to different sizes. A view is returned, not a copy.