Spec-Zone.ru › PyTorch 2.14

torch.tensor_split

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

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

torch.tensor_split(input, sections, dim=0) → Список тензоров

Разбивает input на sections секций вдоль измерения dim. Если input делится на sections без остатка вдоль измерения dim, каждая секция будет иметь одинаковый размер — input.size(dim) / sections. Если input не делится на sections без остатка, размер первых int(input.size(dim) % sections) секций составит int(input.size(dim) / sections) + 1, а размер остальных — int(input.size(dim) / sections).

sections также может быть тензором типа long нулевой размерности.

torch.tensor_split(input, indices, dim=0) → Список тензоров

Разбивает input вдоль измерения dim в каждой из точек, указанных индексами в indices. Например, indices=[2, 3] и dim=0 дадут тензоры input[:2], input[2:3] и input[3:].

indices может быть списком или кортежем целых чисел либо одномерным тензором типа long на CPU.

Параметры:
  • input (Tensor) – тензор для разбиения
  • 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]]))

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.tensor_split.html

Spec-Zone.ru

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