PyTorch torch.expand Function


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

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)

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)

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

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.


Pytorch torch 参考手册Pytorch torch Reference Manual