如何将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
相关产品推荐
相关产品推荐

