Функция stack
Функция stack объединяет последовательность тензоров вдоль новой оси (измерения).
Все тензоры должны иметь одинаковую форму. Первым параметром функция принимает список или кортеж тензоров.
Вторым параметром можно указать ось, вдоль которой будет происходить объединение (по умолчанию dim=0).
Синтаксис
torch.stack(tensors, dim=0)
Пример
Давайте объединим два одномерных тензора вдоль новой нулевой оси:
import torch
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([4, 5, 6])
res = torch.stack([t1, t2])
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3],
[4, 5, 6],
])
Пример
А теперь объединим те же тензоры вдоль новой оси с индексом 1:
import torch
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([4, 5, 6])
res = torch.stack([t1, t2], dim=1)
print(res)
Результат выполнения кода:
tensor([
[1, 4],
[2, 5],
[3, 6],
])
Пример
Объединим три двумерных тензора одинаковой формы:
import torch
t1 = torch.tensor([[1, 2], [3, 4]])
t2 = torch.tensor([[5, 6], [7, 8]])
t3 = torch.tensor([[9, 10], [11, 12]])
res = torch.stack([t1, t2, t3], dim=0)
print(res)
Результат выполнения кода:
tensor([
[[1, 2], [3, 4]],
[[5, 6], [7, 8]],
[[9, 10], [11, 12]],
])