[论文解读] Global Attention Improves Graph Networks Generalization
本文提出低秩全局注意力(LRGA),一种内存与计算效率更高的注意力机制,通过提升图神经网络(GNNs)的泛化能力来增强其性能。通过将LRGA集成到随机图神经网络(RGNN)框架中,该模型实现了与2-Folklore Weisfeiler-Lehman(2-FWL)同构性测试的算法对齐,从而在图分类、回归和链接预测等多个GNN基准测试中达到最先进性能。
This paper advocates incorporating a Low-Rank Global Attention (LRGA) module, a computation and memory efficient variant of the dot-product attention (Vaswani et al., 2017), to Graph Neural Networks (GNNs) for improving their generalization power. To theoretically quantify the generalization properties granted by adding the LRGA module to GNNs, we focus on a specific family of expressive GNNs and show that augmenting it with LRGA provides algorithmic alignment to a powerful graph isomorphism test, namely the 2-Folklore Weisfeiler-Lehman (2-FWL) algorithm. In more detail we: (i) consider the recent Random Graph Neural Network (RGNN) (Sato et al., 2020) framework and prove that it is universal in probability; (ii) show that RGNN augmented with LRGA aligns with 2-FWL update step via polynomial kernels; and (iii) bound the sample complexity of the kernel's feature map when learned with a randomly initialized two-layer MLP. From a practical point of view, augmenting existing GNN layers with LRGA produces state of the art results in current GNN benchmarks. Lastly, we observe that augmenting various GNN architectures with LRGA often closes the performance gap between different models.
研究动机与目标
- 提升图神经网络(GNNs)的泛化能力,尽管其理论表达能力存在局限,但在实践中通常表现良好。
- 解决标准全局注意力在GNN中计算成本过高的问题,其复杂度随图大小呈二次方增长。
- 通过与强大图同构性测试的算法对齐,为GNN中注意力机制的泛化增益提供理论依据。
- 通过实证验证,LRGA能在多种GNN架构和基准测试中提升性能。
提出的方法
- 提出低秩全局注意力(LRGA),一种点积注意力的变体,通过秩-κ近似将复杂度降低至O(κ²|V|)的计算量和O(κ|V|)的内存占用。
- 引入随机图神经网络(RGNN)框架,其中在每次前向传播时重新采样随机特征,证明其在概率上具有通用性。
- 证明在LRGA增强的RGNN中,可通过多项式核学习单项函数,从而与2-FWL图同构性测试实现对齐。
- 为使用随机初始化的两层MLP学习核函数映射时的样本复杂度提供绑定,从而提供泛化保证。
- 在包括OGB和ZINC在内的基准数据集上,对多种GNN架构(GCN、GAT、GraphSage、GatedGCN、GIN)进行LRGA的实证评估。
实验结果
研究问题
- RQ1低秩全局注意力机制是否能超越GNN表达能力的理论极限,提升其泛化能力?
- RQ2通过LRGA增强GNN是否能实现与2-FWL同构性测试的算法对齐,该测试是强于WL测试的图同构性判据?
- RQ3在RGNN框架中,通过LRGA学习2-FWL更新规则的样本复杂度是多少?
- RQ4LRGA是否在多种GNN架构和图学习任务中持续提升性能?
主要发现
- LRGA在所有评估的GNN模型和数据集上均表现更优,常在图分类和回归任务中达到最先进结果。
- 在OGB链接预测基准上,LRGA增强的GCN在ogbl-ppa数据集上达到Hits@100为0.342 ± 0.016,优于Node2vec和DeepWalk。
- 在ogbl-collab数据集上,LRGA + GCN达到Hits@50为0.522 ± 0.007,比第二名的GraphSage高出超过4个百分点。
- 在ogbl-ddi数据集上,LRGA + GCN达到Hits@20为0.623 ± 0.091,显著优于MF和GraphSage。
- 在PATTERN数据集上使用随机特征的消融实验中,LRGA + GIN达到86.765%的准确率,较GIN单独使用提升1.005%。
- LRGA能持续缩小不同GNN架构之间的性能差距,表明其在稳定化和增强泛化能力方面具有关键作用。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。