PyTorch torch.block_diag Function
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)
# 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)
Other Extensions