PyTorch中支持重复索引的张量行累加、平均操作是否有批量实现?
PyTorch 中带重复索引的批量原地实现方案
一、f_sum 的批量原地实现
直接用 PyTorch 内置方法即可实现和循环版本一致的效果,且效率更高:
方法1:使用 index_add_(推荐,专门针对这类场景优化)
# 原地修改v,完成批量累加 v.index_add_(0, indices, w)
方法2:直接索引累加
v[indices] += w
注:当
indices存在重复元素时,PyTorch 会自动对同一索引位置进行多次累加,结果和循环的f_sum完全一致。
二、f_average 的批量原地实现
需要先计算每个索引的出现次数,对 w 做加权处理后再批量累加:
# 计算每个索引的出现次数,确保长度匹配v的第一维度 counts = torch.bincount(indices, minlength=v.size(0)) # 生成每个w对应的权重,广播适配w的形状 weights = 1.0 / counts[indices].unsqueeze(1) # 加权后原地累加至v对应位置 v.index_add_(0, indices, w * weights)
注:
minlength参数避免因索引未出现导致的形状不匹配;由于indices中不存在未出现的索引,counts[indices]不会出现0,无需处理除以0的情况。
内容的提问来源于stack exchange,提问作者Sasha
相关产品推荐
相关产品推荐

