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

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