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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 05:42:13