Spec-Zone.ru › PyTorch 1

torch.tensor_split

torch.tensor_split(input, indices_or_sections, dim=0) → List of Tensors

Разделяет тензор на несколько подтензоров, все из которых являются представлениями input, по размерности dim в соответствии с указанными индексами или количеством разделов, заданными indices_or_sections. Данная функция основана на функции NumPy numpy.array_split().

Параметры:
  • input (Тензор) – тензор для разделения
  • indices_or_sections (Тензор, int или список или кортеж целых чисел) –

    Если indices_or_sections является целым числом n или нульмерным тензором типа long со значением 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 является списком или кортежем целых чисел или одномерным тензором типа long, то 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/1.13/generated/torch.tensor_split.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API