PyTorch torch.triu_indices Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.triu_indicesThis is a function in PyTorch used to generate upper triangular matrix indices. It returns the row and column indices of the upper triangular part (including the diagonal).

Function Definition

torch.triu_indices(row, column, offset=0, dtype=torch.long, device='cpu')

Parameter Description:

  • row: Number of rows
  • column: Number of columns
  • offset: Diagonal offset
  • dtype: Returned data type
  • device: Device

Usage Examples

Example

import torch

# Generate upper triangular indices of a 3x3 matrix
row, col = torch.triu_indices(3, 3)
print("row:", row)
print("col:", col)

The output is:

row: tensor([0, 0, 0, 1, 1, 2])
col: tensor([0, 1, 2, 1, 2, 2])

Example

import torch

# Generate indices and use them for indexing operations
row, col = torch.triu_indices(3, 3, offset=1)

# Create a 3x3 matrix
a = torch.ones(3, 3)

# Use indices to set the upper triangular part
a[row, col] = 0
print(a)

The output is:

tensor([[1., 1., 1.],
        [0., 1., 1.],
        [0., 0., 1.]])

Example

import torch

# Non-square matrix case
row, col = torch.triu_indices(3, 4)
print("row:", row)
print("col:", col)

The output is:

row: tensor([0, 0, 0, 1, 1, 2])
col: tensor([0, 1, 2, 1, 2, 2])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions