PyTorch重复索引下向量化求和与循环求和差异及无循环解决方案
问题原因
你遇到的问题源于PyTorch的索引批量赋值逻辑:result[indices] += values并非逐次累加,而是先计算所有目标位置的最终值再一次性写入。当同一个索引重复出现时,只有最后一次对应的值会被保留,而非累加所有对应值,这就导致了和循环累加的结果差异。
无循环的正确实现
下面是两种高效的向量化解决方案,完全替代循环累加:
方法1:使用torch.scatter_add_(推荐)
scatter_add_是PyTorch原生支持的按索引累加操作,专门适配这类场景,执行效率最高:
import torch indices = torch.LongTensor([0,1,2,1]) values = torch.FloatTensor([1,1,2,2]) result = torch.FloatTensor([0,0,0]) # 按0维度累加,将values对应值加到result的indices指定位置 result.scatter_add_(dim=0, index=indices, src=values) print('scatter_add result:', result) # 输出: tensor([1., 3., 2.])
方法2:使用torch.bincount
如果你的索引均为非负整数,可以用bincount统计每个索引对应的values总和:
import torch indices = torch.LongTensor([0,1,2,1]) values = torch.FloatTensor([1,1,2,2]) # 计算每个索引的累加和,结果长度为最大索引+1 sum_per_index = torch.bincount(indices, weights=values) # 适配目标result的长度 result = torch.zeros(3, dtype=torch.float32) result[:len(sum_per_index)] = sum_per_index print('bincount result:', result) # 输出: tensor([1., 3., 2.])
验证
两种方法的输出均与循环得到的looped_result一致,且都是无循环的向量化操作,在大数据量下的执行效率远高于Python循环。
内容的提问来源于stack exchange,提问作者Ivano Donadi
相关产品推荐
相关产品推荐

