Функция split
Функция split разделяет тензор на несколько подтензоров вдоль заданной оси.
Первым параметром она принимает исходный тензор, вторым - размер каждой части или список размеров,
третьим - ось для разделения (по умолчанию dim=0).
Функция возвращает кортеж из подтензоров.
Синтаксис
torch.split(tensor, split_size_or_sections, dim=0)
Пример
Давайте разделим одномерный тензор из 6 элементов на части по 2 элемента:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6])
res = torch.split(t, 2)
print(res)
Результат выполнения кода:
(tensor([1, 2]), tensor([3, 4]), tensor([5, 6]))
Пример
Разделим двумерный тензор размера 4x3 на части по 2 строки:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12]
])
res = torch.split(t, 2, dim=0)
for i, part in enumerate(res):
print(f"Part {i}: {part}")
Результат выполнения кода:
Part 0: tensor([[1, 2, 3],
[4, 5, 6]])
Part 1: tensor([[7, 8, 9],
[10, 11, 12]])
Пример
Разделим тензор на части неравного размера с помощью списка размеров:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8])
res = torch.split(t, [3, 2, 3])
print(res)
Результат выполнения кода:
(tensor([1, 2, 3]), tensor([4, 5]), tensor([6, 7, 8]))
Пример
Разделим тензор по столбцам (оси 1):
import torch
t = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8]
])
res = torch.split(t, 2, dim=1)
print(res)
Результат выполнения кода:
(tensor([[1, 2],
[5, 6]]),
tensor([[3, 4],
[7, 8]]))
Смотрите также
-
функцию
chunk,
которая разделяет тензор на заданное количество частей -
функцию
tensor_split,
которая разделяет тензор по индексам -
функцию
cat,
которая объединяет тензоры вдоль существующей оси -
функцию
stack,
которая объединяет тензоры вдоль новой оси