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

如何将torch.topk()的TopK值映射到PyTorch空张量对应索引位置

快速实现PyTorch中TopK值到指定索引的批量赋值

不需要用Python循环,直接利用PyTorch的向量化索引操作就能高效完成,底层是C++实现,性能远优于循环遍历。

基础场景(无重复索引)

如果indices中的索引都是唯一的,直接通过索引赋值即可:

import torch

# 假设已定义张量T、K值,以及初始化好的空张量t(如t = torch.zeros_like(T))
value, indices = torch.topk(T, K)
t[indices] = value

这行代码会一次性把value的每个元素对应放到t中indices指定的位置,全程无Python循环,数据量越大效率提升越显著。

进阶场景(存在重复索引)

如果indices里有重复的索引,需要将对应位置的value累加而不是覆盖,可以使用scatter_add_方法:

# 注意需要将indices和value扩展为二维张量(匹配scatter_add_的维度要求)
t.scatter_add_(dim=0, index=indices.unsqueeze(0), src=value.unsqueeze(0))

该方法会把value中的元素累加到t的对应索引位置,避免重复索引导致的覆盖问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 09:37:02