PyTorch torch.topk Function
PyTorch torch Reference Manual
torch.topkIt is a function in PyTorch used to return the largest k elements and their indices.
Function Definition
torch.topk(input, k, dim, largest, sorted, out)
Usage Example
Example
import torch
x = torch.tensor([[3, 1, 2], [6, 4, 5]])
# Return the largest 2 elements
values, indices = torch.topk(x, k=2)
print("The largest 2 values:")
print(values)
print("Corresponding indices:")
print(indices)
x = torch.tensor([[3, 1, 2], [6, 4, 5]])
# Return the largest 2 elements
values, indices = torch.topk(x, k=2)
print("The largest 2 values:")
print(values)
print("Corresponding indices:")
print(indices)
The output result is:
最大的 2 个值:
tensor([[3, 2],
[6, 5]])
对应的索引:
tensor([[0, 2],
[0, 2]])
Other Extensions