PyTorch torch.topk Function


Pytorch torch 参考手册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)

The output result is:

最大的 2 个值:
tensor([[3, 2],
        [6, 5]])
对应的索引:
tensor([[0, 2],
        [0, 2]])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions