Метод reshape
Метод reshape изменяет форму тензора, возвращая новый тензор с теми же данными, но другой размерностью. Метод принимает в качестве аргументов новую форму в виде одного или нескольких целых чисел, либо кортежа. Если новая форма совместима с исходной по общему количеству элементов, то метод возвращает тензор с изменённой формой. В некоторых случаях reshape может вернуть копию данных, если это необходимо для обеспечения непрерывности памяти.
Синтаксис
t.reshape(*shape)
где *shape - это одно или несколько целых чисел, задающих новую форму. Также можно передать кортеж.
Пример
Давайте изменим форму одномерного тензора с 6 элементами в двумерный тензор размера 2x3:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6])
res = t.reshape(2, 3)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[4, 5, 6],
])
Как видно, данные остались теми же, изменилась только форма.
Пример
Использование специального значения -1 для автоматического вычисления размера по одной из осей:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6])
res = t.reshape(2, -1)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[4, 5, 6],
])
Значение -1 указывает PyTorch автоматически вычислить размер этой оси на основе общего количества элементов.
Пример
Изменение формы многомерного тензора с сохранением общего числа элементов:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = t.reshape(3, 2)
print(res)
Результат выполнения кода:
tensor([
[1, 2],
[3, 4],
[5, 6],
])
Общее количество элементов (6) сохранилось, данные переупорядочены в соответствии с новой формой.
Пример
Пример с тензором, для которого reshape создаёт копию данных из-за несовместимости шага (stride):
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
# Транспонируем, чтобы изменить шаг
t = t.transpose(0, 1)
res = t.reshape(2, 3)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[4, 5, 6],
])
В этом случае reshape создаёт копию данных, так как исходный тензор был транспонирован и не является непрерывным.