PyTorch torch.argwhere Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.argwhereis a function in PyTorch used to return the indices of elements that satisfy a condition. It returns the indices of elements whose value is True (non-zero) in the input tensor.

Function Definition

torch.argwhere(input)

Parameters:

  • input(Tensor): The input tensor.

Return value:

  • torch.Tensor: Returns a 2D tensor, each row is an index of an element that satisfies the condition.

Usage Examples

Example

import torch

# Create a tensor
x = torch.tensor([[1, 0, 2],
                  [0, 3, 0],
                  [4, 5, 0]])

print("Original tensor:")
print(x)

# Return the indices of non-zero elements
indices = torch.argwhere(x)

print("nIndices of non-zero elements:")
print(indices)

The output result is:

原始张量:
tensor([[1, 0, 2],
        [0, 3, 0],
        [4, 5, 0]])

非零元素的索引:
tensor([[0, 0],
        [0, 2],
        [1, 1],
        [2, 0],
        [2, 1]])

Example

import torch

# Boolean condition
x = torch.tensor([[True, False, True],
                  [False, True, False],
                  [True, True, False]])

print("Boolean tensor:")
print(x)

indices = torch.argwhere(x)
print("nIndices of True values:")
print(indices)

The output result is:

布尔张量:
tensor([[True, False, True],
        [False, True, False],
        [True, True, False]])

True 值的索引:
tensor([[0, 0],
        [0, 2],
        [1, 1],
        [2, 0],
        [2, 1]])

Example

import torch

# Find elements greater than a certain value
x = torch.randn(3, 4)
threshold = 0

print("Original tensor:")
print(x)

# Find the indices of elements greater than threshold
indices = torch.argwhere(x > threshold)

print(f"nIndices of elements greater than {threshold}:")
print(indices)

# You can also use the nonzero function, with the same effect
indices2 = torch.nonzero(x > threshold)
print("nResult using nonzero:")
print(indices2)

The output result is:

原始张量:
tensor([[-1.2345,  0.5678, -0.8901,  1.2345],
        [ 0.3456, -0.6789,  0.9012, -0.1234],
        [-0.5678,  1.2345, -0.3456,  0.7890]])

大于 0 的元素索引:
tensor([[0, 1],
        [0, 3],
        [1, 0],
        [1, 2],
        [2, 1],
        [2, 3]])

使用 nonzero 的结果:
tensor([[0, 1],
        [0, 3],
        [1, 0],
        [1, 2],
        [2, 1],
        [2, 3]])

Example

import torch

# 1D tensor
x = torch.tensor([1, 0, 0, 4, 0, 5, 0])

indices = torch.argwhere(x)
print("Indices of the 1D tensor:")
print(indices.squeeze())  # Remove the extra dimension

The output result is:

1D张量的索引:
tensor([0, 3, 5])

Example

import torch

# Application: find elements that satisfy a condition and modify them
x = torch.randn(5, 5)

# Find the indices of all elements greater than 0
pos_indices = torch.argwhere(x > 0)

print("Positions in the original tensor greater than 0:")
print(pos_indices)

# Use these indices to modify values
for idx in pos_indices:
    x[idx[0], idx[1]] = 100

print("nModified tensor:")
print(x)

Note:torch.argwhereYestorch.nonzeroan alias of, the two have exactly the same functionality. It returns element indices, each row corresponds to the position of an element satisfying the condition in the original tensor.


Pytorch torch 参考手册PyTorch torch Reference Manual

Other extensions