You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何获取张量中前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]

原理说明

  1. torch.topk(x, n)是PyTorch优化过的原生操作,能以O(m log n)的时间复杂度提取最大的n个元素的索引(m为张量长度),比全量排序更高效。
  2. 对提取出的索引做升序排序,就能得到和示例一致的结果;如果需要保留最大值在原张量中的出现顺序,直接返回top_indices.tolist()即可(示例中会返回[2, 4, 0, 5])。

内容的提问来源于stack exchange,提问作者rosecabbagedragon

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 00:08:17