PyTorch torch.block_diag Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.block_diagis a function in PyTorch used to create a block diagonal matrix. It arranges multiple input tensors as diagonal blocks along the diagonal to form a larger block diagonal matrix.

Function Definition

torch.block_diag(*tensors)

Usage Example

Example

import torch

# Create block diagonal matrix
a = torch.tensor([[1, 2], [3, 4]])
b = torch.tensor([[5, 6, 7]])
c = torch.tensor([[8], [9]])

result = torch.block_diag(a, b, c)
print("Block diagonal matrix:")
print(result)
# tensor([[1, 2, 0, 0, 0],
#         [3, 4, 0, 0, 0],
#         [0, 0, 5, 6, 7],
#         [0, 0, 8, 0, 0],
#         [0, 0, 9, 0, 0]])

# Multiple matrix blocks
m1 = torch.eye(2)
m2 = torch.eye(3)
result2 = torch.block_diag(m1, m2)
print("Block diagonal of two identity matrices:")
print(result2)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions