torch.tensor_split
-
torch.tensor_split(input, indices_or_sections, dim=0) → List of Tensors -
Разделяет тензор на несколько подтензоров, все из которых являются представлениями
input, вдоль размерностиdimв соответствии с указанными индексами или числом разделов, которые задаютсяindices_or_sections. Эта функция основана на функции NumPynumpy.array_split().- Параметры
-
- input (Тензор) – тензор для разделения
-
indices_or_sections (Тензор, int или список или кортеж целых чисел) –
Если
indices_or_sectionsявляется целым числомnили нульмерным тензором с значениемn, тоinputразделяется наnсекций вдоль размерностиdim. Еслиinputделится наnвдоль размерностиdim, то размер каждой секции будет одинаковым,input.size(dim) / n. Еслиinputне делится нацело наn, то размеры первыхint(input.size(dim) % n)секций будут равныint(input.size(dim) / n) + 1, а остальные –int(input.size(dim) / n).Если
indices_or_sectionsявляется списком или кортежем целых чисел, или одномерным тензором, тоinputразделяется вдоль размерностиdimв указанных индексах списка, кортежа или тензора. Например,indices_or_sections=[2, 3]иdim=0приведут к тензорамinput[:2],input[2:3], иinput[3:].Если
indices_or_sectionsявляется тензором, то он должен быть нульмерным или одномерным тензором типа long на CPU. -
dim (int, необязательно) – размерность, вдоль которой нужно разделить тензор. По умолчанию:
0
Пример:
>>> x = torch.arange(8) >>> torch.tensor_split(x, 3) (tensor([0, 1, 2]), tensor([3, 4, 5]), tensor([6, 7])) >>> x = torch.arange(7) >>> torch.tensor_split(x, 3) (tensor([0, 1, 2]), tensor([3, 4]), tensor([5, 6])) >>> torch.tensor_split(x, (1, 6)) (tensor([0]), tensor([1, 2, 3, 4, 5]), tensor([6])) >>> x = torch.arange(14).reshape(2, 7) >>> x tensor([[ 0, 1, 2, 3, 4, 5, 6], [ 7, 8, 9, 10, 11, 12, 13]]) >>> torch.tensor_split(x, 3, dim=1) (tensor([[0, 1, 2], [7, 8, 9]]), tensor([[ 3, 4], [10, 11]]), tensor([[ 5, 6], [12, 13]])) >>> torch.tensor_split(x, (1, 6), dim=1) (tensor([[0], [7]]), tensor([[ 1, 2, 3, 4, 5], [ 8, 9, 10, 11, 12]]), tensor([[ 6], [13]]))
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.tensor_split.html