Метод unbind
Метод unbind удаляет указанное измерение тензора,
возвращая кортеж из тензоров, каждый из которых соответствует
одному срезу по этому измерению. Первым параметром метод
принимает измерение dim (по умолчанию 0), которое будет удалено.
Размерность полученных тензоров будет на единицу меньше исходной.
Синтаксис
tensor.unbind(dim=0)
Пример
Давайте разделим двумерный тензор по строкам (измерение 0):
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
[7, 8, 9],
])
res = t.unbind(dim=0)
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],
[7, 8, 9],
])
res = t.unbind(dim=1)
print(res)
Результат выполнения кода:
(tensor([1, 4, 7]), tensor([2, 5, 8]), tensor([3, 6, 9]))
Мы получили кортеж из трёх тензоров-столбцов.
Пример
Разделим трёхмерный тензор по умолчанию (измерение 0):
import torch
t = torch.tensor([
[
[1, 2],
[3, 4],
],
[
[5, 6],
[7, 8],
],
])
res = t.unbind()
print(res)
Результат выполнения кода:
(tensor([[1, 2], [3, 4]]), tensor([[5, 6], [7, 8]]))
Метод вернул два тензора размерности 2x2.
Пример
Разделим тензор и обработаем каждый срез отдельно:
import torch
t = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
])
for i, row in enumerate(t.unbind()):
print(f"Row {i}: {row.sum().item()}")
Результат выполнения кода:
"Row 0: 10"
"Row 1: 26"