深入理解图神经网络:消息传递、GCN、GraphSAGE、GAT 与工程实践
图神经网络(Graph Neural Network,GNN)最重要的能力,不是“把神经网络用在图上”,而是让每个对象在更新表示时,显式利用它与其他对象之间的关系。
这句话听起来简单,真正落地却涉及一连串问题:邻居信息怎样聚合?为什么 GCN 要做度归一化?GraphSAGE 的“归纳能力”来自哪里?注意力权重能否解释模型?层数加深为什么经常变差?千万节点的图又该怎样训练?
本文从统一的消息传递框架出发,推导四类经典模型,再把重点落到数据划分、大图采样、异构图、常见失败模式和 PyTorch Geometric 实践上。
1. 图数据到底多了什么
一张图通常写作:
其中, 是节点集合, 是边集合。节点特征记作:
是每个节点的特征维度。图还可以包含边特征 、边方向、时间戳、权重以及不同的节点和边类型。
图学习的常见任务分为三类:
- 节点级任务:判断用户是否异常、论文属于哪个领域。
- 边级任务:预测两个人是否会建立联系、两种药物是否发生相互作用。
- 图级任务:预测一个分子的性质、判断一段程序图是否存在漏洞。
与表格或规则网格相比,图的关键差异是:拓扑结构本身也是数据。同样一组节点特征,只要边发生变化,模型面对的样本就已经不同。
卷积神经网络可以利用图像中固定的上下左右邻域;普通图没有统一的节点顺序,每个节点的度数也可能不同。因此,一个合格的图模型至少要解决:
- 如何处理数量不定的邻居;
- 如何保证邻居排列变化时结果不变;
- 如何让信息沿边传播,同时避免高阶邻域带来的噪声和计算爆炸。
2. 统一视角:消息传递神经网络
大部分经典 GNN 都能写成消息传递神经网络(Message Passing Neural Network,MPNN)的形式。
在第 层,邻居 向节点 发送消息:
节点 聚合所有入边消息:
然后更新自己的表示:
这里有三个设计核心:
- 决定“邻居发送什么”;
AGG决定“多条消息怎样合并”;- 决定“节点怎样吸收消息”。
2.1 为什么聚合函数必须与顺序无关
图中的邻居通常是一个集合。即使存储系统改变了边的排列,预测也不应该变化。因此聚合函数应当满足置换不变性:
求和、均值和最大值都满足这一条件,直接把邻居按任意顺序拼接则不满足。
不同聚合器会丢失不同的信息:
mean能表达平均特征,但容易忽略邻居数量;max擅长捕捉显著模式,但会丢失频次;sum可以保留计数信息,配合足够强的 MLP 时表达能力更强。
2.2 层数不等于“真正看见”的距离
理论上,堆叠 层消息传递后,节点可以接收 跳邻域的信息。但“计算图覆盖了 跳”不代表远处信息真的被有效保留。
随着层数增加,至少会同时出现三件事:
- 感受野扩大;
- 节点表示可能越来越相似,即过平滑;
- 大量远距离信息被压入固定维度向量,即过压缩。
所以 GNN 的深度不能机械地类比 CNN。很多节点任务中,两到三层已经是很强的起点。
3. GCN:归一化邻接矩阵上的传播
图卷积网络(Graph Convolutional Network,GCN)是最经典的基线之一。
设邻接矩阵为 。首先加入自环:
再令 为 的度矩阵:
一层 GCN 可以写成:
3.1 公式中的每一项在做什么
加入 是为了让节点在聚合邻居时保留自身信息。否则,没有自环的节点更新后可能完全丢掉自己的原始表示。
左右两侧的度归一化:
会按源节点和目标节点的度共同缩放消息。直观上,高度节点连接很多,如果不归一化,它们及其邻居的数值尺度容易被度数主导。
在节点形式下,GCN 的更新近似为:
3.2 GCN 的优势和边界
GCN 的优点是结构简单、稀疏矩阵实现高效,而且通常是很难绕过的基线。一次传播的主要复杂度与边数近似线性相关,而不是与 相关。
它的局限也很明确:
- 邻居权重由节点度决定,无法根据内容动态选择邻居;
- 更适合相连节点相似的同质性图;
- 全图训练需要保存整张图和中间表示;
- 深层 GCN 容易出现表示退化。
在 PyTorch Geometric 的 GCNConv 中,cached=True 会缓存归一化后的图结构,只适合固定图上的传导式学习。动态图、不同子图批次或邻居采样训练不能盲目开启缓存。
4. GraphSAGE:采样并聚合
GraphSAGE 的核心不是某一个固定卷积公式,而是“采样邻居,再生成节点表示”的归纳式框架。
一种常见写法是:
表示拼接。与为每个节点学习一个独立嵌入不同,GraphSAGE 学的是一套由特征生成表示的聚合函数。只要新节点具有可用特征和邻居,训练结束后加入的新节点也可以被编码。
这就是它的归纳能力来源,而不是“采样”二字本身。
4.1 邻居采样改变了什么
假设每层采样 个邻居,模型有 层,一个种子节点最坏会展开约:
个计算节点。这就是邻居爆炸。采样限制了计算量,却也引入了新的统计问题:
- 高频邻居可能反复出现;
- 低频但关键的邻居可能被漏掉;
- 采样扇出过小会产生较大方差;
- 扇出过大会重新遇到内存瓶颈。
因此,GraphSAGE 通常比全图 GCN 更适合大图,但采样数、层数和批大小必须联合调节。
5. GAT:让模型学习邻居权重
图注意力网络(Graph Attention Network,GAT)不再只依据度数分配权重,而是根据节点表示计算邻居的重要性。
先对节点做线性变换,并计算未归一化注意力:
然后只在节点 的邻域内做 softmax:
最终更新为:
多头注意力会并行计算多组权重,再拼接或求平均。它能让不同头关注不同的邻域模式,并改善优化稳定性。
5.1 注意力不是免费的解释器
GAT 能动态加权邻居,但这不意味着:
- 它在任何图上都优于 GCN;
- 注意力越复杂,表达能力就一定越强;
- 一个较大的注意力权重就是因果解释。
注意力权重受参数化方式、softmax 竞争、特征尺度和邻域构成影响。若要解释预测,应结合遮蔽实验、反事实测试或专门的图解释方法,而不是单看一张注意力热力图。
GAT 还需要为每条边计算注意力分数。在超大图或高度节点上,它的显存和计算成本通常高于简单均值聚合。
6. GIN:从表达能力理解聚合器
图同构网络(Graph Isomorphism Network,GIN)关注一个更理论的问题:消息传递模型能否区分不同的图结构?
其更新形式为:
这里使用求和而不是均值,是因为均值可能把不同的多重集合映射成相同结果。例如, 和 的均值相同,但求和不同。
在合适条件下,GIN 在消息传递 GNN 中达到了与一维 Weisfeiler–Lehman(1-WL)图同构测试相当的区分能力。不过这句话有两个边界:
- 它描述的是模型类别的理论上限,不等于有限数据上的实际精度;
- 1-WL 本身不能区分所有非同构图,因此 GIN 也不是“能识别任意图结构”。
GIN 尤其常见于图级分类。使用求和池化时,模型可以保留与节点数量和局部模式频次有关的信息。
7. 四类模型怎样选
| 模型 | 主要聚合方式 | 主要优势 | 典型限制 | 常见场景 |
|---|---|---|---|---|
| GCN | 度归一化求和 | 简洁、高效、强基线 | 邻居权重固定,偏好同质性 | 中小型固定图、节点分类 |
| GraphSAGE | 采样后 mean/max/LSTM 等 | 归纳学习,适合大图小批量 | 采样方差和邻居爆炸 | 推荐、社交网络、新节点预测 |
| GAT | 邻域注意力加权 | 可按特征动态区分邻居 | 边级计算较贵,权重不等于解释 | 邻居贡献差异明显的任务 |
| GIN | 求和加 MLP | 消息传递框架内表达力强 | 可能放大尺度,对超参敏感 | 分子、程序图等图级任务 |
实践中,更稳妥的顺序通常是:
- 先训练只看节点特征的 MLP;
- 再训练 GCN 或 mean GraphSAGE;
- 证明图结构确实带来增益;
- 最后才根据问题引入注意力、异构关系或更复杂结构编码。
如果简单模型已经没有从边中获得稳定增益,直接换更复杂的 GNN 通常不会修复错误的建图方式。
8. 建图和数据划分比模型名字更重要
GNN 项目最危险的错误,往往发生在模型训练之前。
8.1 节点和边应表达什么
建图时至少要回答:
- 一个节点代表实体、事件,还是实体在某个时间的状态?
- 边是有向还是无向?
- 多条边应该合并、计数,还是保留为不同关系?
- 边是否有发生时间、置信度或权重?
- 不存在边表示“没有关系”,还是“关系尚未被观测”?
例如在交易风控中,把“同一设备登录”与“直接转账”压成相同的无向边,会丢失大量语义。此时异构图或关系类型特定的参数通常更合理。
8.2 先划分,再做可能泄漏信息的图处理
常见的数据泄漏包括:
- 使用未来产生的边预测过去的节点;
- 计算节点统计特征时包含测试期行为;
- 做链接预测时,目标边仍保留在消息传递图中;
- 负样本中混入尚未观测但实际为正的边;
- 同一实体的高度相似副本同时出现在训练集和测试集。
时间敏感任务应优先使用时间切分。链接预测需要明确区分:
- 用于传播消息的边;
- 用于监督训练的正负边;
- 最终评估的边。
训练集上的表现再好,也无法弥补评估协议中的泄漏。
8.3 传导式与归纳式评估
传导式任务允许训练时看到测试节点及其连接,只是不使用测试标签。经典 Cora 节点分类通常属于这一类。
归纳式任务则要求模型面对训练时未出现的节点、子图,甚至全新图。两者回答的问题不同,不能只因为都使用 test_mask 就混为一谈。
9. 用 PyTorch Geometric 训练一个两层 GCN
下面用 Cora 引文网络给出一个可运行的节点分类示例。它的目的不是追逐榜单,而是展示完整且不污染测试集的训练流程。
安装基础依赖:
1 | python -m pip install torch torch-geometric |
如果使用 CUDA,应先依据 PyTorch 和 PyTorch Geometric 官方安装页面选择匹配版本。
1 | import copy |
这段代码中有几个容易被忽略的细节:
- 损失只在
train_mask上计算; - 验证集用于选择模型,测试集只在最后使用一次;
model.eval()会关闭 dropout;- 保存的是最佳参数的深拷贝,而不是仍会继续变化的引用;
cached=True依赖训练、验证和测试始终使用同一张固定图。
严谨实验还应运行多个随机种子,并报告均值和标准差。
10. 大图训练:邻居采样不是简单切批
当整张图无法放入显存时,可以使用 NeighborLoader 从一批种子节点向外逐层采样。
1 | from torch_geometric.loader import NeighborLoader |
若模型有两层,num_neighbors=[15, 10] 表示为不同传播跳数设置采样扇出。采样跳数通常要与模型层数协调,否则可能采了模型用不到的节点,或模型需要的邻域没有被完整展开。
大图系统还应关注:
- CPU 到 GPU 的采样与传输是否成为瓶颈;
- 热点高度节点是否造成批次大小剧烈波动;
- 推理时采用全邻居、分层推理还是同样采样;
- 训练采样分布与线上请求分布是否一致。
11. 图级任务与池化
节点分类为每个节点输出结果,图分类还要把所有节点表示压成一张图的表示:
PyTorch Geometric 中可以使用:
1 | from torch_geometric.nn import global_mean_pool |
图批处理通常不是把节点补齐成相同长度,而是把多张图组合成一张块对角的大图,再用 batch 向量记录每个节点属于哪张图。
池化方式同样会改变模型能表达的信息:
- 均值池化对图大小较稳定,但可能忽略计数;
- 求和池化能保留频次,更符合 GIN 的表达力分析;
- 最大池化突出最显著的局部模式;
- 注意力池化更灵活,但也更容易过拟合。
12. 异构图:不同关系不该共享同一种语义
知识图谱、推荐系统和风控网络经常包含多种节点与边,例如:
1 | (用户)-[购买]->(商品) |
这类图可以写成带类型的关系:
是源节点类型, 是关系类型, 是目标节点类型。异构 GNN 通常会对不同关系使用独立变换,再把同一目标类型收到的多种关系消息进行聚合。
PyTorch Geometric 使用 HeteroData 表示异构图:
1 | from torch_geometric.data import HeteroData |
异构建模要特别警惕:
- 关系类型过多导致参数和显存快速增长;
- 添加反向边时把目标信息泄漏回输入;
- 不同节点类型的特征空间和缺失模式完全不同;
- 某些关系数量巨大,训练时压制了稀有关系。
如果关系方向和类型决定业务语义,先保留它们,再通过消融实验判断是否需要合并,比一开始把图无向化更安全。
13. 两种常被混淆的深层退化
13.1 过平滑:节点越来越像
反复进行邻域平均后,同一连通区域的节点表示可能逐渐趋同,类别之间的可分性下降。这称为过平滑(over-smoothing)。
常见缓解方法包括:
- 控制消息传递层数;
- 使用残差连接或保留初始特征;
- 在不同层之间做 Jumping Knowledge;
- 使用合适的归一化、dropout 或 DropEdge;
- 把特征变换与图传播解耦。
但深层模型变差不能自动归因于过平滑。优化不稳定、训练样本不足和错误建图也可能产生相似现象。
13.2 过压缩:远处信息挤不过来
过压缩(over-squashing)指大量远距离信息必须经过狭窄的拓扑瓶颈,最终被压进固定维度的节点向量。
想象一棵分支很多的树:距离根节点每增加一跳,潜在信息源近似指数增长,但根节点的隐藏维度没有增长。即使这些远处节点都位于理论感受野内,它们的影响也可能被严重压缩。
可能的缓解方向包括:
- 增加虚拟节点或全局连接;
- 对图进行结构重连,缩短关键节点之间的路径;
- 加入位置或结构编码;
- 使用具有全局交互能力的图 Transformer;
- 根据任务减少无关远程消息。
过平滑讨论的是“表示趋同”,过压缩讨论的是“信息穿过瓶颈时丢失”。两者可能同时出现,但不是同一个问题。
14. 同质性、异质性与错误归纳偏置
GCN 等邻域平滑方法隐含了一种偏好:相连节点往往相似。这在论文引用、朋友关系等图中经常成立,却不是普遍规律。
在欺诈检测中,欺诈账户可能大量连接正常账户;在蛋白质网络中,相互作用的节点也可能承担不同功能。这类异质性图上,直接平均邻居可能把有用信号冲淡。
可尝试的方向包括:
- 分离自身表示与邻居表示;
- 区分不同跳数或不同关系的通道;
- 保留方向、符号和边特征;
- 加入结构角色特征,而不只依赖邻居类别相似性;
- 使用针对异质性设计的模型。
最简单的诊断不是立即更换模型,而是计算边两端标签或特征的相关性,并与随机连边基线比较。
15. 一个可信的评估清单
图模型实验至少应回答以下问题。
数据与切分
- 切分是否符合真实部署时间线?
- 测试节点或测试边是否通过预处理泄漏到训练特征?
- 链接预测的监督边是否从消息传递图中移除?
- 新节点、新图和固定图分别采用了什么评估协议?
基线
- 只看节点特征的 MLP 表现如何?
- 树模型或线性模型表现如何?
- 简单标签传播是否已经很强?
- 移除所有边、打乱边或只保留部分关系后会怎样?
指标
- 类别不均衡时不要只看准确率;
- 节点分类可以报告 macro-F1、micro-F1 和每类召回率;
- 稀有正例任务更应关注 PR-AUC;
- 链接预测常用 MRR、Hits@K,并说明负样本生成方式;
- 多个随机种子应报告均值和波动。
消融
至少比较:
- 无图结构;
- 原始图结构;
- 打乱或随机图结构;
- 不同边类型或时间窗口;
- 不同层数与采样扇出。
只有当模型在合理切分和强基线下仍然稳定获益,才能说明“关系结构”确实提供了有效信息。
16. 从问题出发选择模型
可以用下面的顺序快速决策:
图不大、固定、同质性较强
先用两层 GCN。它速度快、实现简单,也便于判断图结构是否有用。
节点不断新增或图太大
先用 mean GraphSAGE 配合邻居采样。重点调试吞吐、扇出和归纳式划分。
不同邻居的贡献确实取决于内容
在 GCN 或 GraphSAGE 基线之上尝试 GAT,同时单独评估额外计算成本。
任务是图分类,并重视局部结构计数
尝试 GIN 与 sum pooling,再与 mean pooling 做消融。
节点和边有明确类型
保留类型并使用异构消息传递。不要为了套用同构 GNN 过早丢掉业务语义。
长距离依赖很重要
先确认问题是否真来自传播距离,再考虑虚拟节点、重连、位置编码或图 Transformer。单纯堆叠更多层通常不是稳健答案。
17. 总结
理解 GNN,可以抓住一条主线:模型通过置换不变的聚合,把图结构转化为可学习的局部消息传递。
GCN 用归一化邻接矩阵传播;GraphSAGE 学习归纳式的采样聚合;GAT 动态计算邻居权重;GIN 用求和与 MLP 提升消息传递框架内的结构区分能力。
但真实项目的上限经常由模型之外的因素决定:边是否表达正确语义、时间切分是否可信、负采样是否合理、邻居采样能否扩展,以及模型究竟遇到过平滑还是过压缩。
如果你正在建立自己的深度学习知识体系,可以继续阅读本站的 CNN 入门教程 和 无监督学习入门。把规则网格、独立样本和关系数据放在一起比较,会更容易理解不同模型的归纳偏置。
参考资料
- Kipf, T. N. and Welling, M. Semi-Supervised Classification with Graph Convolutional Networks.
- Hamilton, W., Ying, Z. and Leskovec, J. Inductive Representation Learning on Large Graphs.
- Veličković, P. et al. Graph Attention Networks.
- Xu, K. et al. How Powerful are Graph Neural Networks?.
- Di Giovanni, F. et al. On Over-Squashing in Message Passing Neural Networks: The Impact of Width, Depth, and Topology.
- Giusti, L. et al. On the Expressive Power of Virtual Nodes in Graph Neural Networks.
- PyTorch Geometric. A Gentle Introduction to PyTorch Geometric.
- PyTorch Geometric. GCNConv API.
- PyTorch Geometric. Scaling GNNs via Neighbor Sampling.
- PyTorch Geometric. Heterogeneous Graph Learning.