PyTorch torch.diagflat Function
Pytorch torch Reference Manual
torch.diagflatIt is a function in PyTorch used to create a diagonal flat matrix. Regardless of whether the input is one-dimensional or multi-dimensional, it will be flattened and used as diagonal elements.
Function Definition
torch.diagflat(input, diagonal=0)
Parameter Description:
input: Input tensor, will be flatteneddiagonal: Diagonal index
Usage Example
Example
import torch
# Create a one-dimensional tensor
x = torch.tensor([1, 2, 3])
# Create a diagonal flat matrix
y = torch.diagflat(x)
print(y)
# Create a one-dimensional tensor
x = torch.tensor([1, 2, 3])
# Create a diagonal flat matrix
y = torch.diagflat(x)
print(y)
The output result is:
tensor([[1, 0, 0],
[0, 2, 0],
[0, 0, 3]])
Example
import torch
# Create a two-dimensional tensor
x = torch.tensor([[1, 2], [3, 4]])
# Will be flattened to [1,2,3,4], creating a 4x4 diagonal matrix
y = torch.diagflat(x)
print(y)
# Create a two-dimensional tensor
x = torch.tensor([[1, 2], [3, 4]])
# Will be flattened to [1,2,3,4], creating a 4x4 diagonal matrix
y = torch.diagflat(x)
print(y)
The output result is:
tensor([[1, 0, 0, 0],
[0, 2, 0, 0],
[0, 0, 3, 0],
[0, 0, 0, 4]])
Other Extensions