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