PyTorch torch.unflatten Function


Pytorch torch 参考手册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 tensor
  • dim: Dimension index to unflatten
  • sizes: 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)

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)

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)

The output result is:

torch.Size([3, 4])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions