Класс PixelShuffle
Класс PixelShuffle выполняет операцию перестановки пикселей,
при которой элементы из каналов переносятся в пространственные размеры.
Это позволяет увеличивать высоту и ширину тензора в upscale_factor раз,
уменьшая при этом количество каналов в upscale_factor² раз.
Данный модуль часто используется в задачах супер-разрешения и повышения
качества изображений.
Синтаксис
torch.nn.PixelShuffle(upscale_factor)
Параметры
Метод принимает один обязательный параметр:
-
upscale_factor(int) - коэффициент увеличения разрешения. Должен быть положительным целым числом.
Пример
Давайте выполним перестановку пикселей с коэффициентом 2 для тензора размером (1, 4, 2, 2):
import torch
import torch.nn as nn
t = torch.tensor([
[
[1, 2],
[3, 4]
],
[
[5, 6],
[7, 8]
],
[
[9, 10],
[11, 12]
],
[
[13, 14],
[15, 16]
]
])
t = t.unsqueeze(0) # добавляем размерность батча
print(t.shape)
pixel_shuffle = nn.PixelShuffle(2)
res = pixel_shuffle(t)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([1, 4, 2, 2])
torch.Size([1, 1, 4, 4])
tensor([[
[
[1, 5, 2, 6],
[9, 13, 10, 14],
[3, 7, 4, 8],
[11, 15, 12, 16]
]
]])
Как видите, тензор размером (1, 4, 2, 2) преобразовался в (1, 1, 4, 4). Количество каналов уменьшилось в 4 раза, а пространственные размеры увеличились в 2 раза.
Пример
Давайте создадим модуль PixelShuffle с коэффициентом 3:
import torch
import torch.nn as nn
pixel_shuffle = nn.PixelShuffle(3)
t = torch.randn(2, 18, 4, 4)
res = pixel_shuffle(t)
print(f"Input shape: {t.shape}")
print(f"Output shape: {res.shape}")
Результат выполнения кода:
"Input shape: torch.Size([2, 18, 4, 4])"
"Output shape: torch.Size([2, 2, 12, 12])"
Количество каналов уменьшилось с 18 до 2 (18 / 3² = 2), а размеры увеличились с 4 до 12 (4 * 3 = 12).
Пример
Применение PixelShuffle в составе последовательной модели:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Conv2d(3, 12, kernel_size=3, padding=1),
nn.PixelShuffle(2),
nn.ReLU()
)
t = torch.randn(1, 3, 64, 64)
res = model(t)
print(f"Output shape: {res.shape}")
Результат выполнения кода:
"Output shape: torch.Size([1, 3, 128, 128])"
Сначала свёрточный слой увеличил количество каналов с 3 до 12,
затем PixelShuffle преобразовал 12 каналов размером 64x64
в 3 канала размером 128x128.
Смотрите также
-
класс
Conv2d,
который выполняет двумерную свёртку -
класс
Upsample,
который увеличивает разрешение тензора -
класс
Unflatten,
который преобразует плоский тензор в многомерный -
класс
ChannelShuffle,
который перемешивает каналы тензора