PyTorch torch.triu Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.triuIt is a function in PyTorch used to extract the upper triangular part of a matrix (including the main diagonal). The lower triangular part will be set to 0.

Function Definition

torch.triu(input, diagonal=0, out=None)

Parameter Description:

  • input: input tensor
  • diagonal: diagonal index, 0 indicates the main diagonal
  • out: output tensor

Usage Examples

Example

import torch

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

# Extract the upper triangular part
y = torch.triu(a)
print(y)

The output result is:

tensor([[1, 2, 3],
        [0, 5, 6],
        [0, 0, 9]])

Example

import torch

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

# Extract the upper triangular part above the first diagonal above the main diagonal
y = torch.triu(a, diagonal=1)
print(y)

The output result is:

tensor([[0, 2, 3],
        [0, 0, 6],
        [0, 0, 0]])

Example

import torch

# Create a non-square matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6]])

# Extract the upper triangular part
y = torch.triu(a)
print(y)

The output result is:

tensor([[1, 2, 3],
        [0, 5, 6]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions