[论文解读] GNNAutoScale: Scalable and Expressive Graph Neural Networks via Historical Embeddings
GNNAutoScale (GAS) 通过使用历史嵌入来剪枝计算图,实现了在大规模图上对任意消息传递图神经网络(GNN)的可扩展训练,无论输入规模如何,GPU 显存使用量均保持恒定。它保持了完整的模型表达能力,并在 ogbn-products 和 Reddit 等大规模基准测试中实现了最先进性能,优于先前的方法,包括 VR-GCN 和 MVS-GNN。
We present GNNAutoScale (GAS), a framework for scaling arbitrary message-passing GNNs to large graphs. GAS prunes entire sub-trees of the computation graph by utilizing historical embeddings from prior training iterations, leading to constant GPU memory consumption in respect to input node size without dropping any data. While existing solutions weaken the expressive power of message passing due to sub-sampling of edges or non-trainable propagations, our approach is provably able to maintain the expressive power of the original GNN. We achieve this by providing approximation error bounds of historical embeddings and show how to tighten them in practice. Empirically, we show that the practical realization of our framework, PyGAS, an easy-to-use extension for PyTorch Geometric, is both fast and memory-efficient, learns expressive node representations, closely resembles the performance of their non-scaling counterparts, and reaches state-of-the-art performance on large-scale graphs.
研究动机与目标
- 为解决由于 GPU 显存限制和邻居爆炸问题,导致在大规模图上训练深层且表达能力强的 GNN 所面临的可扩展性挑战。
- 在实现可扩展训练的同时,保持消息传递 GNN 的完整表达能力,避免因边采样或不可学习传播而导致的性能下降。
- 通过将可扩展性与消息传递机制解耦,实现对多种 GNN 算子(如 GCNII 和 PNA)的通用化应用,超越特定 GNN 架构的限制。
- 提供基于历史嵌入的近似误差的理论边界,并展示实用方法以收紧这些边界。
- 仅使用单个 GPU 并将历史记录存储在 CPU 上,实现对大规模数据集(如 ogbn-products)的全批量 GNN 训练。
提出的方法
- GAS 使用历史嵌入——即前一训练迭代中的节点表示——作为计算图中非小批量节点的内存高效替代方案。
- 它剪枝 GNN 计算图,仅保留当前小批量中的节点及其直接的一跳邻居,从而将 GPU 显存消耗降低至与 GNN 深度无关的恒定水平。
- 历史嵌入在每轮迭代中更新,并用于从非活跃节点传播信息,从而在不存储完整计算图的情况下保留拓扑依赖关系。
- 该框架基于嵌入的陈旧程度和函数的利普希茨连续性,提供了理论近似误差边界,并提出了实用策略以收紧这些边界。
- PyGAS 是 PyTorch Geometric 的扩展,仅需极少代码修改即可实现 GAS,支持与现有 GNN 代码库的无缝集成。
- 该方法与模型设计正交,可应用于任意消息传递 GNN,包括深层(GCNII)和高表达能力(PNA)架构。
实验结果
研究问题
- RQ1我们能否在不进行边采样或牺牲模型表达能力的前提下,将任意消息传递 GNN 扩展到大规模图上?
- RQ2使用历史嵌入引入的理论近似误差是什么?在实践中如何最小化?
- RQ3历史嵌入能否在保持全批量 GNN 表达能力的同时,实现恒定的 GPU 显存使用量?
- RQ4当应用于深层和高表达能力 GNN 时,GAS 框架是否在大规模图基准测试中实现了最先进性能?
- RQ5GAS 能否在无需架构修改的情况下,推广至多种不同的 GNN 架构?
主要发现
- GAS 实现了仅使用单个 GPU 在大规模图(如 ogbn-products,240 万个节点,6190 万个边)上训练全批量 GNN,且 GPU 显存使用量与输入规模无关,保持恒定。
- 在 ogbn-products 数据集上,PNA-GAS 达到 79.91% 的准确率,优于之前最先进方法 GraphSAINT(79.08%),也优于因 OOM 而失败的全批量 PNA。
- 在 Reddit 数据集上,PNA-GAS 达到 97.17% 的准确率,超过 GraphSAINT(97.00%)和 VR-GCN(94.50%),表明其在大规模节点分类任务中具有更优性能。
- GCNII-GAS 在 Reddit 上达到 96.77% 的准确率,优于因 OOM 而失败的全批量 GCNII,表明深层模型可在大规模下成功训练。
- 在 ogbn-products 上,存储历史记录的显存消耗仅为每层约 2GB,可轻松存储于 CPU 内存中,同时训练过程保持高效且可扩展。
- PyGAS 开源实现使开发者仅需极少代码修改,即可在大规模图上训练深层和高表达能力 GNN,性能接近非可扩展基线方法。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。