如何获取张量中前n个最大值的索引?
获取PyTorch张量前n个最大值的索引(含重复值)
要高效实现这个需求,我们可以利用PyTorch原生的torch.topk()函数快速提取目标索引,再按需调整顺序匹配示例输出:
实现代码
import torch def generalized_argmax(x: torch.Tensor, n: int): # 提取前n个最大值的索引(按值从大到小排列) _, top_indices = torch.topk(x, n) # 将索引按升序排序,还原为原张量中的索引顺序 sorted_indices = torch.sort(top_indices)[0] return sorted_indices.tolist()
测试示例
x = torch.tensor([2, 1, 4, 1, 4, 2, 1, 1]) print(generalized_argmax(x, 4)) # 输出: [0, 2, 4, 5]
原理说明
torch.topk(x, n)是PyTorch优化过的原生操作,能以O(m log n)的时间复杂度提取最大的n个元素的索引(m为张量长度),比全量排序更高效。- 对提取出的索引做升序排序,就能得到和示例一致的结果;如果需要保留最大值在原张量中的出现顺序,直接返回
top_indices.tolist()即可(示例中会返回[2, 4, 0, 5])。
内容的提问来源于stack exchange,提问作者rosecabbagedragon
相关产品推荐
相关产品推荐

