Функция reshape
Функция reshape возвращает новый тензор с заданной формой на основе исходного.
Первым параметром функция принимает тензор, вторым - новую форму в виде кортежа или
переменного числа аргументов. Если размерность не указана для одного из измерений
(задаётся как -1), она вычисляется автоматически на основе общего числа элементов.
Синтаксис
torch.reshape(tensor, shape)
Пример
Давайте преобразуем одномерный тензор из 6 элементов в двумерный размером 2x3:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6])
res = torch.reshape(t, (2, 3))
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[4, 5, 6],
])
Пример
Теперь преобразуем тот же тензор в трёхмерный размером 2x1x3 с помощью -1:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6])
res = torch.reshape(t, (2, -1, 3))
print(res)
Результат выполнения кода:
tensor([
[
[1, 2, 3],
],
[
[4, 5, 6],
],
])
Пример
Функцию можно вызывать как метод тензора reshape,
передавая новую форму напрямую:
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],
])
Пример
Если использовать -1 в качестве одного из измерений, форма будет вычислена автоматически. Например, преобразуем одномерный тензор из 12 элементов в четырёхмерный:
import torch
t = torch.arange(12)
res = torch.reshape(t, (2, -1, 2, 3))
print(res.shape)
Результат выполнения кода:
torch.Size([2, 1, 2, 3])