РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
114 of 769 menu

Функция tensor_split

Функция tensor_split разделяет тензор на несколько подтензоров вдоль указанной размерности. Первым параметром функция принимает исходный тензор. Вторым параметром можно передать количество частей или список индексов для разделения. Третьим параметром задаётся ось (по умолчанию разделение происходит по первому измерению).

Синтаксис

torch.tensor_split(tensor, indices_or_sections, [dim])

Пример

Давайте разделим одномерный тензор из 6 элементов на 3 равные части:

import torch t = torch.tensor([1, 2, 3, 4, 5, 6]) res = torch.tensor_split(t, 3) for part in res: print(part)

Результат выполнения кода:

tensor([1, 2]) tensor([3, 4]) tensor([5, 6])

Пример

Если размер не делится нацело, последняя часть будет содержать оставшиеся элементы:

import torch t = torch.tensor([1, 2, 3, 4, 5]) res = torch.tensor_split(t, 3) for part in res: print(part)

Результат выполнения кода:

tensor([1, 2]) tensor([3, 4]) tensor([5])

Пример

Разделим тензор по указанным индексам (после элементов с индексами 2 и 4):

import torch t = torch.tensor([1, 2, 3, 4, 5, 6]) res = torch.tensor_split(t, [2, 4]) for part in res: print(part)

Результат выполнения кода:

tensor([1, 2]) tensor([3, 4]) tensor([5, 6])

Пример

Разделим двумерный тензор по строкам (ось 0):

import torch t = torch.tensor([ [1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], ]) res = torch.tensor_split(t, 2, dim=0) for part in res: print(part)

Результат выполнения кода:

tensor([ [1, 2, 3], [4, 5, 6], ]) tensor([ [7, 8, 9], [10, 11, 12], ])

Пример

Разделим двумерный тензор по столбцам (ось 1):

import torch t = torch.tensor([ [1, 2, 3, 4], [5, 6, 7, 8], ]) res = torch.tensor_split(t, 2, dim=1) for part in res: print(part)

Результат выполнения кода:

tensor([ [1, 2], [5, 6], ]) tensor([ [3, 4], [7, 8], ])

Смотрите также

  • функцию split,
    которая разделяет тензор на части заданного размера
  • функцию chunk,
    которая пытается разделить тензор на равные части
  • функцию cat,
    которая объединяет тензоры вдоль указанной оси
  • функцию unbind,
    которая удаляет указанную ось и возвращает кортеж подтензоров
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить