PyTorch中如何按张量第一列值分组对第二列值求和
PyTorch 张量分组求和实现(无Python循环)
问题场景
现有二维PyTorch张量,第一列为取值范围有限的分组标识,第二列为待计算的数值,需要按第一列的相同值对第二列做分组求和。由于待处理数据量达十亿级,Python层面的for/while循环处理耗时过长,需使用PyTorch原生API实现。
示例输入
import torch val = torch.tensor([[1,233], [1,222], [2,333], [2,3234], [2,3242], [2,3234], [3,234], [3,234], [4,323]])
期望输出
output_val = torch.tensor([[1,455], [2,10043], [3,468], [4,323]])
实现方案
使用PyTorch原生的scatter_add_算子即可实现高性能分组聚合,该算子完全在张量计算后端执行(CPU多线程/GPU CUDA核心),无Python循环开销,处理十亿级数据效率极高。
场景1:第一列分组键为连续整数
如果分组键本身是连续整数(和示例一致),可以直接聚合:
# 提取分组键(转long类型作为索引)和待求和数值 keys = val[:, 0].long() values = val[:, 1] # 初始化聚合结果数组 sum_vals = torch.zeros(keys.max().item() + 1, dtype=values.dtype, device=val.device) # 原地执行分组求和 sum_vals.scatter_add_(0, keys, values) # 整理为要求的输出格式 valid_keys = torch.nonzero(sum_vals).squeeze(1) output_val = torch.stack([valid_keys, sum_vals[valid_keys]], dim=1)
场景2:第一列分组键为非连续整数
如果分组键不是连续值,先用torch.unique将key映射为连续索引再聚合:
keys = val[:, 0].long() values = val[:, 1] # 将非连续key映射为从0开始的连续索引 unique_keys, inverse_idx = torch.unique(keys, return_inverse=True) sum_vals = torch.zeros(unique_keys.shape[0], dtype=values.dtype, device=val.device) sum_vals.scatter_add_(0, inverse_idx, values) # 拼接得到最终结果 output_val = torch.stack([unique_keys, sum_vals], dim=1)
性能说明
- 上述实现全程无Python层面循环,所有计算均由PyTorch后端并行执行,相比Python循环速度可提升数万到数十万倍,十亿级数据通常数分钟内即可处理完成。
- 如果单设备内存/显存无法容纳全量数据,可以分块加载张量,每块按上述逻辑计算局部聚合结果,最后对所有局部结果再做一次二次聚合即可,依然不需要Python循环。
内容的提问来源于stack exchange,提问作者Clock ZHONG
相关产品推荐
相关产品推荐

