PyTorch torch.unflatten Function
PyTorch torch Reference Manual
torch.unflattenis a function in PyTorch used to unflatten dimensions. It expands existing dimensions in a tensor into multiple dimensions to reconstruct tensor shapes.
Function Definition
torch.unflatten(input, dim, sizes)
Parameter Description:
input: Input tensordim: Dimension index to unflattensizes: Tuple of sizes after unflattening
Usage Examples
Example
import torch
# Create a one-dimensional tensor
x = torch.arange(12)
# Unflatten into a 3x4 matrix
y = torch.unflatten(x, dim=0, sizes=(3, 4))
print(y.shape)
print(y)
# Create a one-dimensional tensor
x = torch.arange(12)
# Unflatten into a 3x4 matrix
y = torch.unflatten(x, dim=0, sizes=(3, 4))
print(y.shape)
print(y)
The output result is:
torch.Size([3, 4])
tensor([[ 0, 1, 2, 3],
[ 4, 5, 6, 7],
[ 8, 9, 10, 11]])
Example
import torch
# Create a flattened tensor
x = torch.randn(24)
# Unflatten into a 2x3x4 three-dimensional tensor
y = torch.unflatten(x, dim=0, sizes=(2, 3, 4))
print(y.shape)
# Create a flattened tensor
x = torch.randn(24)
# Unflatten into a 2x3x4 three-dimensional tensor
y = torch.unflatten(x, dim=0, sizes=(2, 3, 4))
print(y.shape)
The output result is:
torch.Size([2, 3, 4])
Example
import torch
# Create a flattened tensor
x = torch.randn(12)
# Can specify the dimension by name
x = x.rename('N')
y = torch.unflatten(x, dim='N', sizes=(3, 4))
print(y.shape)
# Create a flattened tensor
x = torch.randn(12)
# Can specify the dimension by name
x = x.rename('N')
y = torch.unflatten(x, dim='N', sizes=(3, 4))
print(y.shape)
The output result is:
torch.Size([3, 4])
Other Extensions