РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
366 of 769 menu

Класс 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,
    которая применяется для регуляризации и предотвращения переобучения
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить