Класс Unfold
Класс Unfold из модуля torch.nn извлекает скользящие локальные блоки из входного тензора. Он принимает размер окна, шаг и параметры отступов, а затем преобразует входной тензор размерности (N, C, H, W) в тензор размерности (N, C * kernel_height * kernel_width, L), где L - количество блоков. Класс полезен для реализации свёрточных слоёв, операций пулинга и анализа локальных паттернов. Первым параметром передаётся размер окна, вторым - шаг, третьим - отступы.
Синтаксис
torch.nn.Unfold(kernel_size, dilation=1, padding=0, stride=1)
Параметры
Метод Unfold принимает следующие параметры:
-
kernel_size- размер окна (может быть целым числом или кортежем).
Определяет размер извлекаемого блока. -
dilation- шаг между элементами внутри окна (по умолчанию1).
Управляет разрежением окна. -
padding- отступы вокруг входного тензора (по умолчанию0).
Добавляет нулевые отступы перед извлечением. -
stride- шаг скольжения окна (по умолчанию1).
Определяет расстояние между соседними окнами.
Пример
Давайте создадим простой двумерный тензор и извлечём из него окна размером 2x2 с шагом 1:
import torch
from torch.nn import Unfold
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9]
]).float().unsqueeze(0).unsqueeze(0) # shape: (1, 1, 3, 3)
unfold = Unfold(kernel_size=2, stride=1)
res = unfold(t)
print(res)
Результат выполнения кода:
tensor([
[1., 2., 4., 5.],
[2., 3., 5., 6.],
[4., 5., 7., 8.],
[5., 6., 8., 9.]
])
Пример
Рассмотрим извлечение окон с шагом 2, чтобы уменьшить количество блоков:
import torch
from torch.nn import Unfold
t = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12],
[13, 14, 15, 16]
]).float().unsqueeze(0).unsqueeze(0)
unfold = Unfold(kernel_size=2, stride=2)
res = unfold(t)
print(res)
Результат выполнения кода:
tensor([
[ 1., 3., 9., 11.],
[ 2., 4., 10., 12.],
[ 5., 7., 13., 15.],
[ 6., 8., 14., 16.]
])
Пример
Используем отступы для сохранения размера выходных данных:
import torch
from torch.nn import Unfold
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9]
]).float().unsqueeze(0).unsqueeze(0)
unfold = Unfold(kernel_size=3, padding=1, stride=1)
res = unfold(t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 9, 9])