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