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/1.13/generated/torch.flatten.html