Unfold
-
class torch.nn.Unfold(kernel_size, dilation=1, padding=0, stride=1)[source] -
Извлекает скользящие локальные блоки из тензора входных данных с пакетной обработкой.
Рассмотрим тензор с пакетной обработкой
inputформы , где — размерность пакета, — размерность канала, а представляют произвольные пространственные размерности. Эта операция сглаживает каждый скользящий блок размеромkernel_sizeв пространственных измеренияхinputв столбец (т.е., последняя размерность) 3-мерногоoutputтензора формы , где — общее количество значений в каждом блоке (блок имеет пространственных местоположений, каждое из которых содержит вектор с каналами), а — общее количество таких блоков:где формируется пространственными измерениями
input( выше), а — по всем пространственным размерностям.Поэтому индексирование
outputв последней размерности (размерность столбца) даёт все значения в определённом блоке.Аргументы
padding,strideиdilationуказывают, как извлекаются скользящие блоки.-
strideуправляет шагом скользящих блоков. -
paddingуправляет количеством неявных нулевых заполнений по обеим сторонам дляpaddingколичества точек для каждой размерности перед преобразованием формы. -
dilationуправляет интервалом между точками ядра; также известен как алгоритм à trous. Его труднее описать, но эта ссылка имеет хорошее визуальное представление того, что делаетdilation.
- Параметры
-
- kernel_size (int или tuple) — размер скользящих блоков
- dilation (int или tuple, необязательно) — параметр, который управляет шагом элементов в окрестности. По умолчанию: 1
- padding (int или tuple, необязательно) — неявное нулевое заполнение, которое нужно добавить по обеим сторонам входа. По умолчанию: 0
- stride (int или tuple, необязательно) — шаг скользящих блоков в пространственных измерениях входных данных. По умолчанию: 1
- Если
kernel_size,dilation,paddingилиstrideявляется целым числом или кортежем длиной 1, их значения будут дублироваться по всем пространственным размерностям. - В случае двух входных пространственных измерений эта операция иногда называется
im2col.
Примечание
Foldвычисляет каждое комбинированное значение в результирующем большом тензоре, суммируя все значения из всех содержащих блоков.Unfoldизвлекает значения в локальных блоках, копируя их из большого тензора. Таким образом, если блоки перекрываются, они не являются взаимно обратными операциями.В общем, операции сворачивания и развертывания связаны следующим образом. Рассмотрим экземпляры
FoldиUnfoldс одинаковыми параметрами:>>> fold_params = dict(kernel_size=..., dilation=..., padding=..., stride=...) >>> fold = nn.Fold(output_size=..., **fold_params) >>> unfold = nn.Unfold(**fold_params)
Тогда для любого (поддерживаемого)
inputтензора выполняется следующее равенство:fold(unfold(input)) == divisor * input
где
divisor— тензор, зависящий только от формы и типа данныхinput:>>> input_ones = torch.ones(input.shape, dtype=input.dtype) >>> divisor = fold(unfold(input_ones))
Когда тензор
divisorне содержит нулевых элементов, операцииfoldиunfoldявляются взаимно обратными (с точностью до постоянного делителя).Предупреждение
В настоящее время поддерживаются только 4-мерные входные тензоры (тензоры, подобные изображениям, с пакетной обработкой).
- Форма:
-
- Вход:
- Выход: как описано выше
Примеры:
>>> unfold = nn.Unfold(kernel_size=(2, 3)) >>> input = torch.randn(2, 5, 3, 4) >>> output = unfold(input) >>> # each patch contains 30 values (2x3=6 vectors, each of 5 channels) >>> # 4 blocks (2x3 kernels) in total in the 3x4 input >>> output.size() torch.Size([2, 30, 4]) >>> # Convolution is equivalent with Unfold + Matrix Multiplication + Fold (or view to output shape) >>> inp = torch.randn(1, 3, 10, 12) >>> w = torch.randn(2, 3, 4, 5) >>> inp_unf = torch.nn.functional.unfold(inp, (4, 5)) >>> out_unf = inp_unf.transpose(1, 2).matmul(w.view(w.size(0), -1).t()).transpose(1, 2) >>> out = torch.nn.functional.fold(out_unf, (7, 8), (1, 1)) >>> # or equivalently (and avoiding a copy), >>> # out = out_unf.view(1, 2, 7, 8) >>> (torch.nn.functional.conv2d(inp, w) - out).abs().max() tensor(1.9073e-06)
-
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.Unfold.html