Метод flatten
Метод flatten преобразует многомерный тензор в одномерный, объединяя все элементы в одну непрерывную последовательность. Метод принимает необязательные параметры start_dim и end_dim, которые позволяют выравнивать только указанный диапазон осей.
Синтаксис
tensor.flatten(start_dim=0, end_dim=-1)
Параметр start_dim задаёт начальную ось для выравнивания (по умолчанию 0), а end_dim - конечную ось (по умолчанию -1, то есть последняя). Метод возвращает новый тензор с изменённой формой, при этом данные не копируются, если это возможно.
Пример
Давайте создадим двумерный тензор и выровняем его в одномерный:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = t.flatten()
print(res)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5, 6])
Как видим, все элементы тензора объединились в одну строку.
Пример
Давайте рассмотрим трёхмерный тензор и выровняем только две последние оси:
import torch
t = torch.tensor([
[
[1, 2],
[3, 4],
],
[
[5, 6],
[7, 8],
],
])
res = t.flatten(start_dim=1)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([2, 4])
tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
Мы выровняли оси с 1 по последнюю, оставив первую ось неизменной.
Пример
Рассмотрим случай, когда нужно выровнять только центральную часть осей:
import torch
t = torch.tensor([
[
[1, 2, 3],
[4, 5, 6],
],
[
[7, 8, 9],
[10, 11, 12],
],
])
res = t.flatten(start_dim=1, end_dim=2)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([2, 6])
tensor([
[ 1, 2, 3, 4, 5, 6],
[ 7, 8, 9, 10, 11, 12],
])
Выравнивание произошло по двум осям (1 и 2), а нулевая осталась прежней.