PyTorch torch.triu_indices Function
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 rowscolumn: Number of columnsoffset: Diagonal offsetdtype: Returned data typedevice: 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)
# 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)
# 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)
# 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])
Other Extensions