[论文解读] Sequential Aggregation and Rematerialization: Distributed Full-batch Training of Graph Neural Networks on Large Graphs
本文提出顺序聚合与重计算(SAR),一种用于图神经网络(GNNs)的分布式全批量训练方法,通过在反向传播过程中顺序重构并释放计算图组件,显著降低内存使用量。SAR 通过确保每个工作节点的内存使用量与工作节点数量呈线性关系,实现了在大型图(如 ogbn-papers100M 和 ogbn-products)上的训练,相较于 128 个工作节点,内存使用量最高可减少 50%,并结合优化的注意力核函数实现显著加速。
We present the Sequential Aggregation and Rematerialization (SAR) scheme for distributed full-batch training of Graph Neural Networks (GNNs) on large graphs. Large-scale training of GNNs has recently been dominated by sampling-based methods and methods based on non-learnable message passing. SAR on the other hand is a distributed technique that can train any GNN type directly on an entire large graph. The key innovation in SAR is the distributed sequential rematerialization scheme which sequentially re-constructs then frees pieces of the prohibitively large GNN computational graph during the backward pass. This results in excellent memory scaling behavior where the memory consumption per worker goes down linearly with the number of workers, even for densely connected graphs. Using SAR, we report the largest applications of full-batch GNN training to-date, and demonstrate large memory savings as the number of workers increases. We also present a general technique based on kernel fusion and attention-matrix rematerialization to optimize both the runtime and memory efficiency of attention-based models. We show that, coupled with SAR, our optimized attention kernels lead to significant speedups and memory savings in attention-based GNNs.We made the SAR GNN training library publicy available: \\url{https://github.com/IntelLabs/SAR}.
研究动机与目标
- 解决在大规模、高度连接的图上进行 GNN 全批量训练时的内存瓶颈问题,此时计算图变得大到无法存储。
- 实现任意 GNN 架构的可扩展、分布式全批量训练,且不依赖采样或非学习型消息传递机制。
- 通过避免冗余地存储中间张量(尤其是注意力模型),最小化分布式训练中的通信开销。
- 为评估基于采样的和非学习型消息传递 GNN 方法提供一个实用且内存高效的基线。
- 通过消除注意力系数的存储并实时计算,优化基于注意力的 GNN(如 GAT)的性能。
提出的方法
- SAR 使用一种分布式领域并行训练方案,其中输入图被划分到 N 个工作节点上,每个工作节点仅处理其分配的子图。
- 与在前向传播阶段完整存储 GNN 计算图不同,SAR 延迟图的构建,直到反向传播阶段,并分段重构计算图。
- 在反向传播过程中,SAR 顺序地重计算并释放计算图组件,确保任意时刻每个工作节点最多仅存储两个分区。
- 该方法实现了每个工作节点内存使用量为 O(2/N) 的缩放特性,与图的密度无关,通过增加工作节点数量,可实现对任意大规模图的训练。
- 对于基于注意力的模型,SAR 集成了核融合与实时注意力系数计算,避免存储大型注意力矩阵。
- 该方法可推广至任何领域并行训练设置,其中输出依赖于跨工作节点的输入,包括空间并行的 CNN。
实验结果
研究问题
- RQ1是否可以在不依赖采样或非学习型消息传递的前提下,将全批量 GNN 训练扩展到大规模图?
- RQ2如何降低分布式 GNN 训练中的内存使用量,以支持无法装入内存的图的训练?
- RQ3在分布式环境中,于反向传播过程中重计算计算图的通信与内存开销是多少?
- RQ4是否可以对重计算进行优化,以避免在 GAT 模型中存储计算成本较高的中间张量(如注意力系数)?
- RQ5在大规模基准测试中,SAR 与基于采样的和非学习型消息传递的 GNN 方法相比,在内存效率和训练速度方面表现如何?
主要发现
- 在 128 个工作节点上对 ogbn-papers100M 使用 GraphSage 模型训练时,SAR 将每个工作节点的峰值内存消耗最高降低 50%,且内存使用量与 2/N 呈线性关系。
- 在 ogbn-products 上,使用优化的注意力核函数时,SAR 相较于 DGL 的 GAT 实现实现了 2.5 倍的加速,得益于更低的内存压力和实时系数计算。
- 对于 GraphSage,SAR 在通信量更低的情况下实现了与领域并行训练相当的运行时间,同时在 128 个工作节点上将内存使用量减少一半。
- 与 SAR 配套使用的优化注意力核函数(FAK)通过避免存储注意力系数,进一步降低了内存使用量,提升了前向传播速度,且未影响反向传播性能。
- SAR 实现了迄今为止在 ogbn-papers100M 和 ogbn-products 上最大规模的全批量 GNN 训练,证明了在节点数超过一亿的图上训练的可行性。
- 该方法对许多 GNN 变体具有通信避免特性,使得内存节省在运行时开销上近乎免费。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。