Класс ChannelShuffle
Класс ChannelShuffle переставляет каналы входного тензора, разделяя их на группы и перемешивая внутри каждой группы. Модуль принимает на вход тензор размера (N, C, H, W) для 2D данных или (N, C, L) для 1D данных. Первым параметром передаётся количество групп groups, на которые делятся каналы. Входные каналы должны делиться на количество групп нацело. Перестановка выполняется по следующему принципу: входной тензор сначала изменяет форму на (N, groups, C // groups, ...), затем транспонирует размеры групп и каналов, и снова изменяет форму к исходной размерности.
Синтаксис
torch.nn.ChannelShuffle(groups)
Пример
Создадим модуль перестановки каналов с двумя группами и применим его к тензору:
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]
]).float().unsqueeze(0).unsqueeze(0)
channel_shuffle = nn.ChannelShuffle(groups=2)
res = channel_shuffle(t)
print(res)
Результат выполнения кода:
tensor([[[[1., 2., 3., 4.],
[9., 10., 11., 12.],
[5., 6., 7., 8.],
[13., 14., 15., 16.]]]])
Пример
Применим перестановку каналов к случайному тензору с размерностью 4 канала:
import torch
import torch.nn as nn
torch.manual_seed(0)
t = torch.randn(1, 4, 2, 2)
channel_shuffle = nn.ChannelShuffle(groups=2)
res = channel_shuffle(t)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([1, 4, 2, 2])
tensor([[[[ 1.5410, -0.2934],
[-2.1788, 0.5684]],
[[-1.0845, -1.3986],
[ 0.4033, 0.8380]],
[[ 0.5427, -0.5392],
[ 0.0396, -1.5542]],
[[-0.3414, -0.1195],
[ 1.0063, 0.6774]]]])
Пример
Используем модуль в составе последовательной модели:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Conv2d(8, 8, kernel_size=3, padding=1),
nn.ChannelShuffle(groups=4),
nn.ReLU()
)
torch.manual_seed(0)
t = torch.randn(1, 8, 5, 5)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 8, 5, 5])
Смотрите также
-
класс
PixelShuffle,
который переставляет элементы из каналов в пространственные измерения -
класс
Conv2d,
который выполняет двумерную свертку с входным тензором -
класс
GroupNorm,
который применяет групповую нормализацию к входному тензору -
класс
Unflatten,
который изменяет форму тензора, восстанавливая размерности