Функция F.conv_transpose2d
Функция F.conv_transpose2d выполняет операцию транспонированной двумерной свертки (также известной как деконволюция) над входным тензором. Эта операция является обратной по отношению к обычной свертке и позволяет увеличивать пространственное разрешение входных данных. Функция принимает входной тензор, тензор весов и опциональные параметры, такие как смещение (bias), шаг (stride), дополнение (padding) и коэффициент расширения (dilation).
Основное отличие от обычной свертки заключается в том, что транспонированная свертка преобразует входной тензор с низким разрешением в выходной тензор с более высоким разрешением. Это достигается за счет того, что каждый элемент входного тензора умножается на ядро свертки, и результаты накладываются друг на друга с заданным шагом. Данная операция широко применяется в архитектурах генеративных нейронных сетей, таких как GAN и вариационные автокодировщики, а также в задачах семантической сегментации.
Синтаксис
torch.nn.functional.conv_transpose2d(
input,
weight,
bias=None,
stride=1,
padding=0,
output_padding=0,
groups=1,
dilation=1,
padding_mode='zeros'
)
Параметры функции:
-
input- входной тензор формы (batch_size, in_channels, height, width) -
weight- тензор весов формы (in_channels, out_channels, kernel_height, kernel_width) -
bias- опциональный тензор смещения формы (out_channels) -
stride- шаг свертки (число или кортеж), по умолчанию 1 -
padding- дополнение входного тензора (число или кортеж), по умолчанию 0 -
output_padding- дополнительное дополнение выхода (число или кортеж), по умолчанию 0 -
groups- количество групп для групповой свертки, по умолчанию 1 -
dilation- коэффициент расширения ядра (число или кортеж), по умолчанию 1 -
padding_mode- режим дополнения: 'zeros', 'reflect', 'replicate' или 'circular'
Пример
Давайте выполним базовую операцию транспонированной свертки с ядром 3x3:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
input_t = torch.randn(1, 1, 4, 4)
weight_t = torch.randn(1, 1, 3, 3)
res = F.conv_transpose2d(input_t, weight_t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 1, 6, 6])
Выходной тензор имеет размер 6x6, так как транспонированная свертка увеличивает пространственное разрешение.
Пример
Изменим шаг свертки для увеличения разрешения выходного тензора:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
input_t = torch.randn(1, 1, 3, 3)
weight_t = torch.randn(1, 1, 2, 2)
res = F.conv_transpose2d(input_t, weight_t, stride=2)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 1, 6, 6])
При шаге 2 выходной тензор имеет размер 6x6, что в два раза больше входного по каждому измерению.
Пример
Добавим смещение (bias) и используем большее количество выходных каналов:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
input_t = torch.randn(2, 3, 5, 5)
weight_t = torch.randn(3, 6, 3, 3)
bias_t = torch.randn(6)
res = F.conv_transpose2d(input_t, weight_t, bias_t, stride=1, padding=1)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 6, 5, 5])
Выходной тензор сохраняет пространственные размеры (5x5) благодаря дополнению padding=1, при этом количество каналов увеличивается с 3 до 6.
Пример
Рассмотрим использование параметра output_padding для точного контроля размера выхода:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
input_t = torch.randn(1, 1, 3, 3)
weight_t = torch.randn(1, 1, 2, 2)
res1 = F.conv_transpose2d(input_t, weight_t, stride=2, output_padding=0)
res2 = F.conv_transpose2d(input_t, weight_t, stride=2, output_padding=1)
print(res1.shape)
print(res2.shape)
Результат выполнения кода:
torch.Size([1, 1, 6, 6])
torch.Size([1, 1, 7, 7])
Параметр output_padding позволяет компенсировать неоднозначность размера выхода и точно контролировать итоговую размерность.
Смотрите также
-
функцию
conv2d,
которая выполняет обычную двумерную свертку -
функцию
conv1d,
которая выполняет одномерную свертку -
функцию
interpolate,
которая изменяет размер тензора с помощью интерполяции -
функцию
pixel_shuffle,
которая переупорядочивает элементы для увеличения разрешения