Метод squeeze
Метод squeeze возвращает новый тензор с удалёнными размерностями,
длина которых равна 1. Если размерность не указана, то удаляются все
одиночные размерности. В противном случае удаляется только указанная
размерность. Исходный тензор остаётся неизменным.
Синтаксис
t.squeeze([dim])
Пример
Давайте создадим тензор с размерностью 1 по первому измерению
и удалим её с помощью метода squeeze:
import torch
t = torch.tensor([
[1, 2, 3, 4, 5]
])
print('shape:', t.shape)
res = t.squeeze()
print('res shape:', res.shape)
print(res)
Результат выполнения кода:
shape: torch.Size([1, 5])
res shape: torch.Size([5])
tensor([1, 2, 3, 4, 5])
Пример
Теперь удалим размерность по индексу 0, используя параметр dim:
import torch
t = torch.tensor([
[1, 2, 3, 4, 5]
])
print('shape:', t.shape)
res = t.squeeze(dim=0)
print('res shape:', res.shape)
print(res)
Результат выполнения кода:
shape: torch.Size([1, 5])
res shape: torch.Size([5])
tensor([1, 2, 3, 4, 5])
Пример
Если размерность по указанному индексу не равна 1, то тензор
останется без изменений:
import torch
t = torch.tensor([
[1, 2, 3, 4, 5]
])
print('shape:', t.shape)
res = t.squeeze(dim=1)
print('res shape:', res.shape)
print(res)
Результат выполнения кода:
shape: torch.Size([1, 5])
res shape: torch.Size([1, 5])
tensor([[1, 2, 3, 4, 5]])