Класс Unflatten
Класс Unflatten является частью модуля torch.nn и предназначен для преобразования плоского тензора в тензор с заданной размерностью. Он выполняет операцию, обратную Flatten. Первым параметром указывается размерность, которую необходимо развернуть, вторым - новая форма или кортеж размеров.
Синтаксис
torch.nn.Unflatten(dim, unflattened_size)
Параметры
Параметр dim определяет, какая размерность будет разворачиваться. Он может быть задан как целое число или строка. Параметр unflattened_size задаёт новую форму для указанной размерности. Это может быть кортеж или список целых чисел.
Пример
Создадим слой Unflatten для преобразования одномерного тензора в двумерный:
import torch
import torch.nn as nn
layer = nn.Unflatten(dim=1, unflattened_size=(2, 3))
t = torch.randn(4, 6)
res = layer(t)
print(res.shape)
Результат выполнения кода:
torch.Size([4, 2, 3])
Пример
Используем Unflatten в составе последовательной модели для восстановления формы после Flatten:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Flatten(),
nn.Linear(12, 12),
nn.Unflatten(dim=1, unflattened_size=(3, 4))
)
t = torch.randn(2, 3, 4)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 3, 4])
Пример
Работа с именованными размерностями с использованием параметра dim в виде строки:
import torch
import torch.nn as nn
t = torch.randn(8, 16, 32)
layer = nn.Unflatten(dim='features', unflattened_size=(4, 8))
res = layer(t)
print(res.shape)
Результат выполнения кода:
torch.Size([8, 4, 8, 32])