Атрибут shape
Атрибут shape класса Tensor возвращает размерность тензора в виде кортежа целых чисел. Он используется для получения информации о количестве элементов по каждой оси тензора. Данный атрибут доступен только для чтения и не изменяет сам тензор.
Синтаксис
tensor.shape
Пример
Давайте создадим одномерный тензор и получим его размерность:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
print(t.shape)
Результат выполнения кода:
torch.Size([5])
Пример
Теперь создадим двумерный тензор и посмотрим его форму:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
print(t.shape)
Результат выполнения кода:
torch.Size([2, 3])
Пример
Атрибут shape часто используется в сочетании с другими атрибутами и методами. Например, можно получить общее количество элементов через умножение всех размерностей:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
res = t.shape[0] * t.shape[1]
print(res)
Результат выполнения кода:
9
Однако для вычисления общего количества элементов в тензоре лучше использовать метод numel, который возвращает тот же результат, но делает это более эффективно.
Пример
Атрибут shape также используется для получения размерности тензора при динамических операциях. Например, можно создать новый тензор с такой же формой:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
t_new = torch.zeros(t.shape)
print(t_new)
Результат выполнения кода:
tensor([
[0., 0., 0.],
[0., 0., 0.],
])