Класс 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,
который выполняет эффективное увеличение разрешения из каналов