torch.flatten
-
torch.flatten(input, start_dim=0, end_dim=-1) → Tensor -
Расплющивает
input, переформировав его в одномерный тензор. Еслиstart_dimилиend_dimпереданы, только измерения, начиная сstart_dimи заканчиваяend_dim, сплющиваются. Порядок элементов вinputне изменяется.В отличие от NumPy’s flatten, который всегда копирует данные входных данных, эта функция может возвращать исходный объект, представление или копию. Если ни одно измерение не сплющено, то возвращается исходный объект
input. В противном случае, если входные данные могут быть представлены как сплющенная форма, возвращается это представление. И только если входные данные не могут быть представлены как сплющенная форма, данные входных данных копируются. См.torch.Tensor.view()для подробностей о том, когда будет возвращено представление.Примечание
Расплющивание тензора нулевой размерности вернёт одномерное представление.
- Параметры
Пример:
>>> t = torch.tensor([[[1, 2], ... [3, 4]], ... [[5, 6], ... [7, 8]]]) >>> torch.flatten(t) tensor([1, 2, 3, 4, 5, 6, 7, 8]) >>> torch.flatten(t, start_dim=1) tensor([[1, 2, 3, 4], [5, 6, 7, 8]])
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.flatten.html