Метод split
Метод split разделяет тензор на несколько подтензоров вдоль указанной размерности.
Первым параметром метод принимает размер каждой части (целое число) или количество частей (список или кортеж).
Вторым параметром передается размерность, по которой производится разбиение (по умолчанию dim=0).
Метод возвращает кортеж тензоров.
Синтаксис
t.split(split_size, dim=0)
Пример
Давайте разобьем одномерный тензор на части по 2 элемента:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6])
res = t.split(2)
print(res)
Результат выполнения кода:
(tensor([1, 2]), tensor([3, 4]), tensor([5, 6]))
Пример
Если размер тензора не делится нацело на размер части, последняя часть будет меньше:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = t.split(2)
print(res)
Результат выполнения кода:
(tensor([1, 2]), tensor([3, 4]), tensor([5]))
Пример
Разобьем двумерный тензор по строкам (размерность 0) на части по 2 строки:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
[10, 11, 12],
])
res = t.split(2, dim=0)
print(res)
Результат выполнения кода:
(tensor([
[1, 2, 3],
[4, 5, 6],
]), tensor([
[7, 8, 9],
[10, 11, 12],
]))
Пример
Разобьем двумерный тензор по столбцам (размерность 1) на части по 2 столбца:
import torch
t = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
res = t.split(2, dim=1)
print(res)
Результат выполнения кода:
(tensor([
[1, 2],
[5, 6],
]), tensor([
[3, 4],
[7, 8],
]))
Пример
Вместо размера части можно передать список или кортеж с размерами каждой части:
import torch
t = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9])
res = t.split([2, 3, 4])
print(res)
Результат выполнения кода:
(tensor([1, 2]), tensor([3, 4, 5]), tensor([6, 7, 8, 9]))