Функция 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],
])