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

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