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

Класс Upsample

Класс Upsample из модуля torch.nn предназначен для увеличения пространственного разрешения входных данных. Слой применяет интерполяцию для масштабирования изображений, карт признаков или других многомерных данных. Основные параметры: scale_factor - коэффициент масштабирования или size - целевой размер выходного тензора, а также mode - метод интерполяции ('nearest', 'linear', 'bilinear', 'bicubic', 'trilinear') и align_corners - способ выравнивания угловых пикселей.

Синтаксис

torch.nn.Upsample( size=None, scale_factor=None, mode='nearest', align_corners=None )

Входной тензор должен иметь форму (batch, channels, height, width) для 2D данных или (batch, channels, depth, height, width) для 3D данных.

Пример с масштабированием

Давайте создадим слой Upsample с коэффициентом масштабирования 2 и режимом 'nearest':

import torch import torch.nn as nn upsample = nn.Upsample( scale_factor=2, mode='nearest' ) t = torch.randn(1, 1, 2, 2) res = upsample(t) print(res.shape)

Результат выполнения кода:

torch.Size([1, 1, 4, 4])

Пример с указанием целевого размера

Увеличим изображение до конкретного размера с помощью билинейной интерполяции:

import torch import torch.nn as nn upsample = nn.Upsample( size=(8, 8), mode='bilinear', align_corners=True ) t = torch.randn(1, 3, 4, 4) res = upsample(t) print(res.shape)

Результат выполнения кода:

torch.Size([1, 3, 8, 8])

Пример работы с 3D данными

Применим трилинейную интерполяцию для объемных данных:

import torch import torch.nn as nn upsample = nn.Upsample( scale_factor=2, mode='trilinear', align_corners=False ) t = torch.randn(1, 1, 2, 2, 2) res = upsample(t) print(res.shape)

Результат выполнения кода:

torch.Size([1, 1, 4, 4, 4])

Пример использования в модели

Встроим слой Upsample в последовательную модель для увеличения карт признаков:

import torch import torch.nn as nn model = nn.Sequential( nn.Conv2d(3, 16, 3, padding=1), nn.ReLU(), nn.Upsample( scale_factor=2, mode='bilinear', align_corners=True ), nn.Conv2d(16, 3, 3, padding=1) ) t = torch.randn(1, 3, 32, 32) res = model(t) print(res.shape)

Результат выполнения кода:

torch.Size([1, 3, 64, 64])

Смотрите также

  • класс ConvTranspose2d,
    который выполняет транспонированную свертку для увеличения разрешения
  • класс UpsamplingBilinear2d,
    который специализируется на билинейной интерполяции для 2D данных
  • класс UpsamplingNearest2d,
    который использует интерполяцию ближайшего соседа для 2D данных
  • класс PixelShuffle,
    который выполняет эффективное увеличение разрешения из каналов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить