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

能否用PyTorch的scatter或gather操作复现torch_geometric中的聚合函数?

能否用PyTorch的scatter或gather操作复现torch_geometric中的聚合函数?

嘿,这个问题我之前折腾PyTorch和PyG的时候也琢磨过,答案是完全可以用纯PyTorch的scatter系列操作来复现torch_geometric里的大部分核心聚合函数!先给你理清楚逻辑:gather更多是按索引从张量里提取元素,而实现组聚合的核心其实是scatter的各种归约变体,比如scatter_add_、scatter_max_这些,正好对应PyG里的各类聚合逻辑,下面给你逐个对应举例:

  • 求和聚合(对应PyG的SumAggregation)
    用torch.scatter_add_就能直接实现,它会把源张量中对应索引位置的元素累加到目标张量里。举个实际的代码例子:

    # 假设x是形状为[N, F]的特征张量,index是形状为[N]的分组索引(每个特征属于哪个组)
    group_num = index.max() + 1
    # 初始化输出张量,形状为[group_num, F]
    sum_out = torch.zeros(group_num, x.size(1), device=x.device)
    # 把x按index分组求和
    sum_out.scatter_add_(0, index.unsqueeze(1).expand(-1, x.size(1)), x)
    

    这段代码的效果和PyG里直接用SumAggregation完全一致。

  • 最大值/最小值聚合(对应PyG的MaxAggregation/MinAggregation)
    对应使用torch.scatter_max_和torch.scatter_min_,这两个函数会返回聚合后的结果以及对应的原始索引(我们只需要结果部分):

    # 最大值聚合
    max_out, _ = torch.scatter_max(x, 0, index.unsqueeze(1).expand(-1, x.size(1)))
    # 最小值聚合
    min_out, _ = torch.scatter_min(x, 0, index.unsqueeze(1).expand(-1, x.size(1)))
    
  • 均值聚合(对应PyG的MeanAggregation)
    均值需要先做求和,再除以每个组的元素数量,同样用scatter_add_配合计数实现:

    group_num = index.max() + 1
    # 第一步:统计每个组的元素个数
    count = torch.zeros(group_num, device=x.device)
    count.scatter_add_(0, index, torch.ones_like(index, dtype=x.dtype))
    # 第二步:计算每个组的特征总和
    sum_out = torch.zeros(group_num, x.size(1), device=x.device)
    sum_out.scatter_add_(0, index.unsqueeze(1).expand(-1, x.size(1)), x)
    # 第三步:计算均值
    mean_out = sum_out / count.unsqueeze(1)
    

    这里可以额外加个小细节:如果存在空分组(count为0),可以给分母加个极小值1e-10防止除零错误,这和PyG内部的处理逻辑也对齐。

其实PyG里的Aggregation模块很多都是基于这些scatter操作封装的,只是额外加了一些便捷处理(比如可学习权重、归一化、多维度适配等),核心的组归约逻辑完全可以用纯PyTorch的scatter系列操作复现。至于gather,它的定位是按索引提取元素,更适合“分散取数”的场景,而非“多元素归约到一组”的聚合需求,所以一般不用它来实现聚合。

备注:内容来源于stack exchange,提问作者daqh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 19:29:33