torch.split
-
torch.split(tensor, split_size_or_sections, dim=0)[source] -
Разделяет тензор на части. Каждая часть является представлением исходного тензора.
Если
split_size_or_sectionsявляется целым типом, тоtensorбудет разделен на части равного размера (если возможно). Последняя часть будет меньше, если размер тензора по указанному измерениюdimне делится наsplit_size.Если
split_size_or_sectionsявляется списком, тоtensorбудет разделен наlen(split_size_or_sections)частей с размерами вdimв соответствии сsplit_size_or_sections.- Параметры
- Возвращаемое значение
Пример:
>>> a = torch.arange(10).reshape(5, 2) >>> a tensor([[0, 1], [2, 3], [4, 5], [6, 7], [8, 9]]) >>> torch.split(a, 2) (tensor([[0, 1], [2, 3]]), tensor([[4, 5], [6, 7]]), tensor([[8, 9]])) >>> torch.split(a, [1, 4]) (tensor([[0, 1]]), tensor([[2, 3], [4, 5], [6, 7], [8, 9]]))
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.split.html