Метод unflatten
Метод unflatten класса Tensor применяется к тензору для восстановления его исходной формы,
которая была изменена при помощи метода flatten или других операций.
Первый параметр определяет измерение (или несколько измерений), которое требуется развернуть.
Второй параметр указывает на новый размер, который может быть задан в виде кортежа или списка.
Метод полезен при работе с пакетами данных, последовательностями или многомерными признаками, когда необходимо вернуть тензор к исходной структуре после свёртки или линейных преобразований.
Важно отметить, что общее количество элементов в тензоре после операции unflatten должно совпадать
с количеством элементов до применения метода. Это условие гарантирует отсутствие потери данных.
Синтаксис
tensor.unflatten(dim, sizes)
Пример
Давайте создадим тензор размерности (2, 3, 4) и применим к нему метод flatten, чтобы получить одномерный тензор из 24 элементов, а затем восстановим форму с помощью unflatten:
import torch
t = torch.tensor([[[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]],
[[13, 14, 15, 16],
[17, 18, 19, 20],
[21, 22, 23, 24]]])
print("Original shape:", t.shape)
flat_t = t.flatten()
print("Flattened shape:", flat_t.shape)
res = flat_t.unflatten(0, (2, 3, 4))
print("Unflattened shape:", res.shape)
Результат выполнения кода:
Original shape: torch.Size([2, 3, 4])
Flattened shape: torch.Size([24])
Unflattened shape: torch.Size([2, 3, 4])
Пример
Теперь рассмотрим ситуацию, когда нужно восстановить только одно измерение из нескольких, например, после операции с пакетом последовательностей:
import torch
t = torch.tensor([[1, 2, 3, 4, 5, 6],
[7, 8, 9, 10, 11, 12]])
print("Original shape:", t.shape)
# flatten первых двух измерений в одно
flat_t = t.flatten(0, 1)
print("Flattened shape:", flat_t.shape)
# восстановление до формы (2, 6)
res = flat_t.unflatten(0, (2, 6))
print("Unflattened shape:", res.shape)
Результат выполнения кода:
Original shape: torch.Size([2, 6])
Flattened shape: torch.Size([12])
Unflattened shape: torch.Size([2, 6])
Пример
Метод unflatten также позволяет восстановить несколько измерений одновременно, передав размеры для каждого из них:
import torch
t = torch.tensor([[[1, 2, 3, 4],
[5, 6, 7, 8]],
[[9, 10, 11, 12],
[13, 14, 15, 16]]])
print("Original shape:", t.shape)
flat_t = t.flatten()
print("Flattened shape:", flat_t.shape)
res = flat_t.unflatten(0, (2, 2, 4))
print("Unflattened shape:", res.shape)
Результат выполнения кода:
Original shape: torch.Size([2, 2, 4])
Flattened shape: torch.Size([16])
Unflattened shape: torch.Size([2, 2, 4])
Пример
В этом примере мы покажем, как использовать unflatten для восстановления формы тензора, полученного после применения flatten с указанием диапазона размерностей:
import torch
t = torch.tensor([[[1, 2, 3],
[4, 5, 6]],
[[7, 8, 9],
[10, 11, 12]],
[[13, 14, 15],
[16, 17, 18]]])
print("Original shape:", t.shape)
# flatten последних двух измерений
flat_t = t.flatten(1, 2)
print("Flattened shape:", flat_t.shape)
# восстановление размеров по измерению 1 на (2, 3)
res = flat_t.unflatten(1, (2, 3))
print("Unflattened shape:", res.shape)
Результат выполнения кода:
Original shape: torch.Size([3, 2, 3])
Flattened shape: torch.Size([3, 6])
Unflattened shape: torch.Size([3, 2, 3])