PyTorch torch.argwhere Function
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)
# 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)
# 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)
# 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
# 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)
# 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.
Other extensions