Функция unbind
Функция unbind разделяет тензор на несколько тензоров меньшей размерности.
Она удаляет указанную размерность (ось) и возвращает кортеж из срезов вдоль этой оси.
Первый параметр - это входной тензор, второй параметр dim задаёт ось для разбиения (по умолчанию 0).
Результат - кортеж тензоров, каждый из которых соответствует одному элементу вдоль указанной оси.
Синтаксис
torch.unbind(tensor, [dim=0])
Пример
Давайте разобьём двумерный тензор на отдельные строки (по умолчанию ось 0):
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
res = torch.unbind(t)
print(res)
Результат выполнения кода:
(tensor([1, 2, 3]), tensor([4, 5, 6]), tensor([7, 8, 9]))
Пример
Теперь разобьём тензор по столбцам, указав ось 1:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = torch.unbind(t, dim=1)
print(res)
Результат выполнения кода:
(tensor([1, 4]), tensor([2, 5]), tensor([3, 6]))
Пример
Работа с трёхмерным тензором: разбиение по глубине (ось 0):
import torch
t = torch.tensor([
[[1, 2], [3, 4]],
[[5, 6], [7, 8]],
])
res = torch.unbind(t, dim=0)
print(res)
Результат выполнения кода:
(tensor([[1, 2], [3, 4]]), tensor([[5, 6], [7, 8]]))
Пример
Разбиение одномерного тензора даёт кортеж из скалярных тензоров:
import torch
t = torch.tensor([10, 20, 30, 40])
res = torch.unbind(t)
print(res)
Результат выполнения кода:
(tensor(10), tensor(20), tensor(30), tensor(40))