Перестановка осей в PyTorch
Для перестановки осей целиком
используют метод permute.
В аргументе перечисляют старые
индексы осей в том порядке, в каком
они должны идти в новой форме.
Число элементов не меняется.
Возьмём трёхмерный тензор: два слоя, в каждом таблица из двух рядов по два числа. Создадим его и переставим оси так, чтобы первой шла бывшая третья:
import torch
stack = torch.tensor([[[1, 2], [3, 4]], [[5, 6], [7, 8]]])
print(stack.shape) # выведет torch.Size([2, 2, 2])
turned = stack.permute(2, 0, 1)
print(turned.shape) # выведет torch.Size([2, 2, 2])
Сами значения при перестановке
перекладываются по новым осям.
Выведем тензор после
permute:
import torch
stack = torch.tensor([[[1, 2], [3, 4]], [[5, 6], [7, 8]]])
turned = stack.permute(2, 0, 1)
print(turned)
Соберите трёхмерный тензор
формы (2, 3, 1) из
последовательных целых чисел
и выведите форму после перестановки
осей в порядке 2, 0, 1.
Соберите трёхмерный тензор
формы (1, 2, 2) с
любыми значениями и выведите
форму после перестановки осей
в порядке 1, 2, 0.
Соберите трёхмерный тензор
из двух таблиц 2×2
и выведите форму после перестановки
осей в порядке 0, 2, 1.