Следите за новинками
в нашем Telegram канале. Жми, чтобы подписаться:)
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 для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить