F.interpolate
Функция interpolate из модуля torch.nn.functional используется для изменения размерности (масштабирования) многомерных тензоров, таких как изображения, карты признаков или объемные данные. Она поддерживает различные режимы интерполяции, включая ближайшего соседа, билинейную, трилинейную и другие. Основными параметрами являются входной тензор, целевой размер или коэффициент масштабирования, а также режим интерполяции. Функция часто применяется в задачах компьютерного зрения для приведения тензоров к единому размеру, например, перед подачей в сверточные слои.
Синтаксис
torch.nn.functional.interpolate(
input,
size=None,
scale_factor=None,
mode='nearest',
align_corners=None,
recompute_scale_factor=None,
antialias=False
)
Основные параметры:
input- входной тензор размерности(N, C, *), где*- пространственные размеры;size- целевой пространственный размер (кортеж или список), например,(h, w);scale_factor- множитель масштабирования (число или кортеж), при указании этого параметра размер вычисляется автоматически;mode- метод интерполяции: 'nearest', 'linear', 'bilinear', 'bicubic', 'trilinear', 'area' (по умолчанию 'nearest');align_corners- логический флаг, определяющий, выравниваются ли углы пикселей, актуален для билинейной и бикубической интерполяции;antialias- приTrueвключает сглаживание для билинейной и бикубической интерполяции (доступно в последних версиях PyTorch).
Пример
Выполним простейшее масштабирование одноканального изображения (тензора) размером 2x2 до размера 4x4 методом ближайшего соседа:
import torch
import torch.nn.functional as F
t = torch.tensor([
[1.0, 2.0],
[3.0, 4.0]
]).unsqueeze(0).unsqueeze(0) # (1, 1, 2, 2)
res = F.interpolate(t, size=(4, 4), mode='nearest')
print(res)
Результат выполнения кода:
tensor([[[[1., 1., 2., 2.],
[1., 1., 2., 2.],
[3., 3., 4., 4.],
[3., 3., 4., 4.]]]])
Пример
Масштабирование цветного изображения (3 канала) с использованием билинейной интерполяции и коэффициента масштабирования 2:
import torch
import torch.nn.functional as F
t = torch.randn(1, 3, 4, 4) # (batch, channels, height, width)
res = F.interpolate(t, scale_factor=2.0, mode='bilinear', align_corners=False)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 3, 8, 8])
Пример
Использование режима 'area' для понижения разрешения (усреднение пикселей):
import torch
import torch.nn.functional as F
t = torch.tensor([
[1.0, 2.0, 3.0, 4.0],
[5.0, 6.0, 7.0, 8.0],
[9.0, 10.0, 11.0, 12.0],
[13.0, 14.0, 15.0, 16.0]
]).unsqueeze(0).unsqueeze(0) # (1, 1, 4, 4)
res = F.interpolate(t, size=(2, 2), mode='area')
print(res)
Результат выполнения кода:
tensor([[[[3.5000, 5.5000],
[11.5000, 13.5000]]]])
Пример
Масштабирование трехмерного объема (например, видео или медицинские данные) с трилинейной интерполяцией:
import torch
import torch.nn.functional as F
t = torch.randn(1, 1, 4, 4, 4) # (batch, channels, depth, height, width)
res = F.interpolate(t, size=(8, 8, 8), mode='trilinear', align_corners=False)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 1, 8, 8, 8])
Смотрите также
-
функцию
conv2d,
которая применяет двумерную свертку к входному тензору -
функцию
max_pool2d,
которая выполняет операцию подвыборки с нахождением максимума -
функцию
avg_pool2d,
которая выполняет усредняющую подвыборку -
функцию
adaptive_avg_pool2d,
которая приводит тензор к заданному размеру с помощью адаптивного усреднения