Функция block_diag
Функция block_diag создает блочно-диагональную матрицу из переданных тензоров. Первым параметром функция принимает переменное количество тензоров. Тензоры могут быть одномерными или двумерными, при этом одномерные тензоры интерпретируются как матрицы размером 1xN.
Синтаксис
torch.block_diag(*tensors)
Пример
Давайте создадим блочно-диагональную матрицу из двух двумерных тензоров:
import torch
t1 = torch.tensor([[1, 2], [3, 4]])
t2 = torch.tensor([[5, 6], [7, 8]])
res = torch.block_diag(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 0, 0],
[3, 4, 0, 0],
[0, 0, 5, 6],
[0, 0, 7, 8],
])
Пример
Давайте создадим блочно-диагональную матрицу из одномерных тензоров:
import torch
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([4, 5])
res = torch.block_diag(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3, 0, 0],
[0, 0, 0, 4, 5],
])
Пример
Давайте создадим блочно-диагональную матрицу из тензоров разных размерностей:
import torch
t1 = torch.tensor([1, 2])
t2 = torch.tensor([[3, 4], [5, 6]])
t3 = torch.tensor([7, 8, 9])
res = torch.block_diag(t1, t2, t3)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 0, 0, 0, 0, 0],
[0, 0, 3, 4, 0, 0, 0],
[0, 0, 5, 6, 0, 0, 0],
[0, 0, 0, 0, 7, 8, 9],
])
Пример
Давайте создадим блочно-диагональную матрицу из тензоров с фиксированным зерном для воспроизводимости:
import torch
torch.manual_seed(0)
t1 = torch.randint(1, 10, (2, 2))
t2 = torch.randint(1, 10, (2, 3))
t3 = torch.randint(1, 10, (3, 2))
res = torch.block_diag(t1, t2, t3)
print(res)
Результат выполнения кода:
tensor([
[4, 9, 0, 0, 0, 0, 0],
[3, 0, 0, 0, 0, 0, 0],
[0, 0, 3, 5, 2, 0, 0],
[0, 0, 7, 6, 8, 0, 0],
[0, 0, 0, 0, 0, 6, 7],
[0, 0, 0, 0, 0, 0, 0],
[0, 0, 0, 0, 0, 6, 2],
])
Пример
Давайте используем функцию block_diag для создания матрицы из трех тензоров с разными размерами:
import torch
t1 = torch.tensor([[1, 2, 3], [4, 5, 6]])
t2 = torch.tensor([[7], [8]])
t3 = torch.tensor([[9, 10]])
res = torch.block_diag(t1, t2, t3)
print(res)
Результат выполнения кода:
tensor([
[1, 2, 3, 0, 0, 0],
[4, 5, 6, 0, 0, 0],
[0, 0, 0, 7, 0, 0],
[0, 0, 0, 8, 0, 0],
[0, 0, 0, 0, 9, 10],
])