PyTorch torch.diagflat Function


Pytorch torch 参考手册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 flattened
  • diagonal: 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)

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)

The output result is:

tensor([[1, 0, 0, 0],
        [0, 2, 0, 0],
        [0, 0, 3, 0],
        [0, 0, 0, 4]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions