Класс Flatten
Класс Flatten из модуля torch.nn преобразует входной тензор в одномерный,
выравнивая все измерения, начиная с указанной позиции. Первым параметром
конструктор принимает start_dim (по умолчанию 1) - индекс измерения,
с которого начинается выравнивание. Вторым параметром передаётся
end_dim (по умолчанию -1) - индекс измерения, которым заканчивается
выравнивание. Первое измерение (batch-измерение) по умолчанию не выравнивается,
что сохраняет размер батча.
Синтаксис
torch.nn.Flatten(start_dim=1, end_dim=-1)
Пример
Давайте создадим слой Flatten с параметрами по умолчанию
и применим его к трехмерному тензору:
import torch
import torch.nn as nn
flatten = nn.Flatten()
t = torch.tensor([
[[1, 2], [3, 4]],
[[5, 6], [7, 8]],
])
res = flatten(t)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
Размерность тензора изменилась с 2×2×2 на 2×4,
при этом первое измерение (размер батча) осталось нетронутым.
Пример
Давайте укажем параметр start_dim, равный 0,
чтобы выровнять все измерения включая batch:
import torch
import torch.nn as nn
flatten = nn.Flatten(start_dim=0)
t = torch.tensor([
[[1, 2], [3, 4]],
[[5, 6], [7, 8]],
])
res = flatten(t)
print(res)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5, 6, 7, 8])
Пример
Давайте укажем параметры start_dim и end_dim,
чтобы выровнять только часть измерений:
import torch
import torch.nn as nn
flatten = nn.Flatten(start_dim=1, end_dim=2)
t = torch.tensor([
[[[1, 2], [3, 4]]],
[[[5, 6], [7, 8]]],
])
res = flatten(t)
print(res)
print(res.shape)
Результат выполнения кода:
tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
torch.Size([2, 4])
Исходный тензор имел размер 2×1×2×2. После выравнивания
измерений с 1 по 2 получилась размерность 2×4.
Пример
Давайте используем Flatten внутри последовательной модели:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Conv2d(3, 16, kernel_size=3),
nn.Flatten(),
nn.Linear(16 * 6 * 6, 10),
)
t = torch.randn(4, 3, 8, 8)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([4, 10])
Слой Flatten преобразовал выход сверточного слоя
размером 4×16×6×6 в тензор размером 4×576,
который затем был подан на полносвязный слой.
Смотрите также
-
класс
Unflatten,
который выполняет обратную операцию - преобразует выровненный тензор в многомерный -
функцию
Linear,
которая применяет линейное преобразование к входным данным -
функцию
Conv2d,
которая выполняет двумерную свертку и часто используется передFlatten -
функцию
Dropout,
которая применяется для регуляризации и предотвращения переобучения