Функция F.pixel_shuffle
Функция F.pixel_shuffle выполняет операцию, обратную F.pixel_unshuffle.
Она принимает тензор формы (N, C * r^2, H, W) и преобразует его в тензор формы
(N, C, H * r, W * r), где r - это коэффициент уменьшения глубины.
Другими словами, элементы из каналов переставляются в пространственные измерения,
увеличивая высоту и ширину в r раз и уменьшая количество каналов в r^2 раз.
Первым параметром функция принимает входной тензор input,
вторым - коэффициент уменьшения глубины upscale_factor.
Синтаксис
torch.nn.functional.pixel_shuffle(input, upscale_factor)
Пример
Давайте преобразуем тензор размером (1, 4, 2, 2) с коэффициентом 2:
import torch
import torch.nn.functional as F
t = torch.tensor([
[
[[1, 2], [3, 4]],
[[5, 6], [7, 8]],
[[9, 10], [11, 12]],
[[13, 14], [15, 16]]
]
], dtype=torch.float)
res = F.pixel_shuffle(t, 2)
print(res)
Результат выполнения кода:
tensor([
[
[[1., 2., 5., 6.],
[3., 4., 7., 8.],
[9., 10., 13., 14.],
[11., 12., 15., 16.]]
]
])
Пример
Используем коэффициент 3 для тензора с 18 каналами:
import torch
import torch.nn.functional as F
t = torch.randn(1, 18, 2, 2)
res = F.pixel_shuffle(t, 3)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 2, 6, 6])
Пример
Используем pixel_shuffle в составе нейросети для повышения разрешения:
import torch
import torch.nn as nn
import torch.nn.functional as F
class UpsampleNet(nn.Module):
def __init__(self):
super(UpsampleNet, self).__init__()
self.conv = nn.Conv2d(3, 12, 3, padding=1)
def forward(self, x):
x = self.conv(x)
x = F.pixel_shuffle(x, 2)
return x
model = UpsampleNet()
t = torch.randn(1, 3, 4, 4)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 3, 8, 8])
Смотрите также
-
функцию
pixel_unshuffle,
которая выполняет обратную операцию перестановки пикселей -
функцию
interpolate,
которая изменяет размер тензора с использованием различных методов интерполяции -
функцию
conv2d,
которая применяет двумерную свертку к входному тензору -
функцию
unfold,
которая извлекает скользящие блоки из входного тензора