尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

GNN入门学习路线:从GCN到GAT,再到动态图与异构图

GNN入门学习路线:从GCN到GAT,再到动态图与异构图 GNN 图神经网络这几年已经成了推荐系统、分子性质预测、知识图谱、风控、交通流量预测里绕不开的技术方向。很多初学者第一次接触 GNN 会有点懵输入既不是图片也不是文本而是一张带点、带边、带属性的图传统卷积层和全连接层没法直接套用。这篇文章打算按一条真正能走通的学习路线把 GNN 从基础图卷积网络 GCN到图注意力机制 GAT再到动态图、异构图这些经典方向完整串一遍。适合刚入门图神经网络的学生和工程师也适合想系统整理知识体系的开发者。最值得先记住一句话GNN 无论怎么变核心都在于让每个节点通过聚合邻居信息学到更合理的节点表示。1. 先搞清楚GNN 解决的是哪一类问题1.1 图数据为什么不能直接套用普通神经网络普通深度学习处理的数据结构比较固定。图片是二维像素网格文本是词序列语音是一维时间序列。图不一样一张图里的节点数量可变每个节点的邻居数量也可变节点和节点之间不是独立同分布的。社交网络里的用户、推荐系统里的商品、分子里的原子、知识图谱里的实体这些对象天然存在连接关系。直接把邻接矩阵铺平当成全连接网络的输入问题很多邻接矩阵非常大真实图通常极度稀疏直接展开浪费严重。节点顺序一变矩阵就完全变了同一个图换个编号方式模型输出就不同。没有参数共享学到的权重只在固定大小的图上有效没法泛化到更大或更小的图。GNN 的思路是在“节点本地”做特征变换再反复聚合邻居信息。这样节点编号怎么变聚合逻辑不变模型天然对节点顺序不敏感。这是它能处理图数据的根本原因。1.2 GNN 最常见的三类任务节点分类、链接预测、图分类GNN 能处理的任务很多但绝大多数入门案例都落在三类上。节点分类给每个节点预测一个标签。典型场景是论文引用网络里判断论文属于哪个研究方向社交网络里判断用户是否存在风险蛋白质网络里判断蛋白质功能。这类任务最常用也是理解 GNN 最好的入口。链接预测预测两个节点之间有没有边或者未来会不会出现边。推荐系统里预测用户会不会购买某个商品知识图谱里预测缺失的关系风控里预测两个账号是否存在关联都属于这一类。图分类把整张图映射成一个标签或数值。分子性质预测是最典型场景一个分子是一张图原子是节点化学键是边模型学习整张图的性质。除了这三类还有图生成、图匹配、社区发现等方向但核心原理是一样的先把节点变成有意义的表示再根据任务设计输出层和损失函数。1.3 核心概念节点表示、邻居聚合、感受野理解 GNN建议先抓住三个词。第一个是“节点表示”。每个节点初始有一组特征可能是用户年龄、交易金额、论文关键词向量也可能是通过 embedding 学出来的。GNN 的每一层都会更新这组表示让节点表示包含越来越丰富的邻居信息。第二个是“邻居聚合”。每次更新节点表示时模型先看这个节点的邻居把邻居的信息按某种方式合并起来再和节点自身信息结合。这个“某种方式”正是不同 GNN 模型的差别所在GCN 用度归一化固定权重GAT 用注意力动态算权重GraphSAGE 用采样加拼接。第三个是“感受野”。对图数据来说GNN 每走一层节点能感知的范围就往外扩一跳。层数越多一个节点能参考的邻居范围越广。听起来层数越多越好但实际图中层数加深会带来过平滑问题后面单独说。2. 消息传递框架所有 GNN 的共同底座2.1 图的常用表示特征矩阵、邻接矩阵、边索引先统一记法。一个图记作 G (V, E)V 是节点集合E 是边集合。工程里常用三样东西表示图特征矩阵 X形状是 [N, D]N 是节点数D 是每个节点的特征维度。邻接矩阵 A形状是 [N, N]A[i][j] 表示节点 i 和节点 j 之间是否有边。标签 Y形状根据任务而定节点分类就是 [N, C]C 是类别数。实际代码里几乎不会真的用 N x N 的稠密矩阵而是用稀疏格式。PyTorch Geometric 里常见的是 edge_index形状为 [2, E]第一行存每条边的源节点第二行存目标节点一列代表一条边。为什么不直接用邻接矩阵因为真实图动辄百万节点稠密矩阵根本存不下即使存得下大部分位置都是 0计算量也浪费在无意义的乘法上。稀疏格式只记录有边的位置训练和推理都更高效。2.2 消息传递的三大步骤消息、聚合、更新所有 GNN 层都可以拆成三步消息、聚合、更新。第一步计算消息。每个节点把自己的表示变成一条“消息”发给邻居。最简单的情况是消息就是节点表示本身复杂模型可以加上可学习变换。第二步聚合。节点把收到的一堆邻居消息合并成一个向量。常见聚合方式有求和、求平均、取最大值也有更复杂的注意力加权求和。第三步更新。节点把聚合结果和自身上一层的表示合并通过一层线性变换和非线性激活函数得到新表示。合并方式可以是相加、拼接也可以是更复杂的门控。用伪代码表示就是for each node v: messages [transform(h_u) for u in neighbors(v)] aggregated aggregate(messages) h_v_new update(h_v, aggregated)理解这个框架后再看各种 GNN 其实都是在回答同一个问题消息怎么算邻居怎么聚合自己和邻居怎么合并。2.3 层数和感受野的关系GNN 一层只看一跳邻居两层能间接看到两跳邻居三层看到三跳。层数决定了节点能利用多远的信息。在节点分类里如果节点标签主要依赖局部结构两层通常够了。链路预测需要判断两个节点之间的关联往往需要更多跳数但也不是越多越好。超过一定层数每个节点的表示会越来越相似分类边界会糊掉这个现象叫过平滑。实际调参时不要一上来就堆深度。先用一层或两层模型跑通再根据结果考虑要不要加深。深度带来的收益通常没有清洗数据、优化特征来得快。2.4 从朴素 GNN 到现代 GNN为什么放弃迭代收敛最早的图神经网络受循环神经网络影响通过反复迭代让节点表示达到不动点一套参数在图上反复更新直到收敛。这种方式理论上优雅实际效率很低要等迭代收敛才能得到输出训练也慢。后来大家发现与其递归迭代到稳定不如直接叠有限层每层用不同参数最后加一个任务相关输出层。这样训练成为标准的端到端 supervised learning各种优化器、dropout、batch normalization 都能直接用。现在提到 GNN默认都是这种有限层的堆叠模型。3. GCN 图卷积网络从公式到工程直觉3.1 GCN 的一次传播到底做了什么GCN全称 Graph Convolutional Network是 2017 年 Kipf 和 Welling 提出的模型也是绝大多数人入门的第一个 GNN。它的传播公式可以写成H^{(l1)} σ( A_hat · H^{(l)} · W^{(l)} )其中 H^{(l)} 是第 l 层所有节点的表示矩阵W^{(l)} 是可学习参数A_hat 是处理过的邻接矩阵σ 是激活函数通常用 ReLU。只看公式会觉得有点抽象。拆开看GCN 一层其实就是对每个节点做了一次“带权重的邻居求和”先对邻居特征做线性变换再把邻居信息累加起来最后过激活函数。这个权重不是网络学出来的而是由图的度结构预先算好的。3.2 自环和归一化矩阵的含义A_hat 不是原始邻接矩阵而是做两步处理后的结果。第一步加自环写成 A I。I 是单位矩阵相当于给每个节点加一条指向自己的边。为什么必须加因为如果聚合时只看邻居更新后的表示就完全不包含节点自身信息变成了“邻居平均值”这会导致同一个节点在不同局部结构里信息丢失。加上自环等于聚合范围变成“自己加邻居”。第二步度归一化写成 D^{-1/2} A_hat D^{-1/2}。D 是度矩阵。直接对邻居求和有问题某个节点如果有几百个邻居聚出来的数值天然很大另一个节点只有两个邻居数值很小。特征尺度不一致训练会不稳定。归一化后信息传递和节点度数解耦模型更容易训练。这也是我建议手写一遍 GCN 实现的原因。真正写过一次 A_hat 的计算就会明白这两个操作不是公式摆设而是让模型能稳定训练的工程细节。3.3 为什么两层 GCN 是经典配置GCN 论文里的经典结构是两层例如Z softmax(A_hat · ReLU(A_hat · X · W0) · W1)第一层把原始特征投影到隐藏维度过 ReLU做一次邻居聚合第二层再聚合一次输出每个类别上的概率。两层 GCN 的实践中表现通常不错原因在于大多数图数据的标签信息都能在局部邻域内找到。对引用网络、社交网络这类数据一个节点属于哪个类别往往看它的直接邻居和二次邻居就够判断了。再加层数不仅训练更慢过平滑风险也明显上升。3.4 GCN 的边界有向图、深度、邻居数量GCN 在无向图上效果稳定但这不意味着所有场景都适用。有向图要额外处理入边、出边和双向边。多关系图比如知识图谱每条边还有类型GCN 无法直接区分。深度加深过平滑这是另一个问题。邻居数量差异极大时即使做了归一化也可能出现信息被少数高影响节点主导的情况。实际项目里不要默认 GCN 就是最优它更像一个靠谱的 baseline。先跑一个 GCN拿到一个可复现的分数再考虑更复杂的模型。3.5 GraphSAGE邻居采样让大规模图可以训练GCN 做全图训练时每一层卷积都要把所有节点计算一遍。图一大内存和计算量都扛不住。GraphSAGE 的思路是每次前向传播时不要所有邻居都参与聚合而是随机采样固定数量的邻居。比如第一层采样 25 个邻居第二层采样 10 个计算量从“跟节点度数相关”变成“跟采样数相关”可控性大幅提升。它的聚合方式也更多样除了平均还有 LSTM 聚合和池化聚合。LSTM 聚合把邻居顺序打乱后按序列处理表达能力更强但计算更重。工程上如果资源有限平均聚合已经够用。GraphSAGE 适合大规模图的另一个原因是它天然支持归纳式学习模型学到的是聚合逻辑新节点来了只要有邻居特征就能直接预测不必重新训练整个图。3.6 图扩散卷积 GDC扩大感受野的一种改进GCN 每次只聚合一跳邻居两层也就两跳。有些任务需要更长距离的信息但直接加深 GCN 又容易过平滑。图扩散卷积的思路是把“在图上游走”的扩散过程显式建模出来用扩散矩阵替换原始邻接矩阵。扩散矩阵可以让信息不只是流向直接邻居而是按步长权重流向更远的节点。这样即使模型层数不多也能利用更大范围的图结构信息。这类方法在有监督分类中不一定总能大幅提升但在部分图数据上尤其当类别标签和长程结构强相关时会比普通 GCN 有更稳定的表现。学习时可以先了解原理不一定要立刻用它替换主模型。4. GAT 图注意力机制让节点自己判断邻居权重4.1 从固定权重到可学习的注意力系数GCN 聚合邻居时权重由图的度结构决定同一个节点的所有邻居在归一化后的权重是固定的不会因为特征差异而变化。但现实里邻居重要性差异很大判断一个用户是不是风险用户某个异常交易邻居可能比一百个普通好友更有信号价值。GAT全称 Graph Attention Network2018 年提出核心改动就是让模型自己学习每个邻居的权重。它先计算节点 i 和邻居 j 之间的注意力系数e_{ij} LeakyReLU(a^T · [W h_i || W h_j])这里 W 是特征变换矩阵h_i 和 h_j 是两端的节点表示|| 表示拼接a 是可学习向量。算出原始注意力分数后在邻居集合上做 softmax 归一化得到 alpha_{ij}表示节点 j 对节点 i 的相对重要程度。聚合时每个邻居的表示乘上对应的 alpha_{ij} 再加起来h_i_new σ( Σ_j alpha_{ij} · W h_j )这样每个节点对不同的邻居给予不同权重表达能力比 GCN 更强尤其在邻居质量参差不齐的数据上。4.2 多头注意力到底加了什么为了让注意力更稳定GAT 借鉴 Transformer 的思路使用多头机制。设置 K 个头每个头独立计算注意力系数独立聚合邻居信息最后把所有头的输出拼接起来或者求平均。多头的好处是降低注意力随机性。单头注意力可能过拟合到某些固定的邻居关系上多头相当于多个视角每个头关注不同的结构模式总体更鲁棒。代价也很直接计算量大约涨 K 倍训练时间明显变长。我见过不少人一上来就设 8 个头小规模数据直接过拟合。建议从 4 个头开始配合 dropout观察验证集效果再调整。4.3 GCN 和 GAT 如何选型很多人纠结该用哪个其实可以按下面这张表快速判断对比项GCNGAT邻居权重来源度归一化固定注意力动态学习表达能力中等更强训练开销低高对邻居噪声的容忍度一般更好小数据上的过拟合风险中等偏高推荐入门顺序第一个试作为进阶对比如果数据量大、邻居重要性差异明显比如风控、推荐GAT 更值得尝试。如果只是课程项目或快速验证GCN 更稳妥训练快参数少调起来省时间。4.4 注意力机制的常见陷阱注意力不是万能的有几种情况要留意。第一注意力分布退化。如果训练后注意力权重接近均匀说明模型没有真正利用邻居差异。可能是特征信息量不足也可能是数据本身邻居确实都差不多。第二注意力过拟合。GAT 在小图上很容易过度依赖某些训练节点测试时换个网络结构就崩。缓解方法是加 dropout或者把多头输出做平均而不是拼接。第三计算复杂度。GAT 要算每对邻居之间的注意力系数边数多时内存开销大。超大图不建议直接全图训练要配合邻居采样或考虑简化变体。5. 动态图和异构图真实系统里的两种“复杂图”5.1 动态图边和节点随时间变化到目前讨论的图都是静态图边固定节点固定特征固定。真实系统不是这样。社交网络里好友关系不断新增交易网络里每条交易都有时间戳交通路网上车流随时间变化。这些图叫动态图。动态图建模要回答的问题有两类一是当前时刻节点状态是什么二是一段时间后两个节点是否会发生连接或交互。前者类似节点分类的时序版本后者类似带时间的链接预测。5.2 快照建模和事件流建模动态图主要有两种建模路线。一种是时间快照路线叫离散时间动态图。把时间按天、周、小时切成段每个时间段内的边组成一个静态图分别跑 GNN再对时间维度建模。可以用循环神经网络处理快照序列也可以用 Transformer 捕捉跨段时间依赖。这种做法实现简单适合数据本身按批次更新、不需要实时预测的场景。另一种是事件流路线叫连续时间动态图。每条边带时间戳模型需要按时间顺序处理事件。常见做法包括时间编码、时序邻居采样、记忆模块更新。比如一个用户在某时刻关注了某个博主这个事件会更新用户的状态后续预测要基于最新状态。风控和实时推荐往往需要这种建模方式。从工程落地角度如果数据按日更新快照方法完全够用如果事件实时到达必须用连续时间建模。很多人拿日报数据硬套连续时间模型训练慢收益却不大。5.3 异构图多类型节点和多类型关系异构图指节点类型和边类型都不止一种。学术网络里有作者、论文、会议三种节点作者和论文之间有“写”关系论文和会议之间有“发表”关系论文和论文之间有“引用”关系。电商场景里用户、商品、品牌加上点击、购买、收藏行为也是标准异构图。异构图的难点在于不同类型节点特征维度可能不同不同关系语义不同。聚合成一个模型时不能直接把所有邻居混在一起求和必须区分类型否则关系语义就没了。5.4 关系聚合与元路径处理异构图有两种主流方式。第一种是关系型聚合。每种边类型配一个线性变换矩阵节点分别按不同类型关系聚合邻居再把结果合并。这样“用户点击商品”和“用户购买商品”可以学到不同的变换语义不会混在一起。问题是关系多时参数膨胀容易过拟合需要强正则。第二种是元路径。元路径是一条跨类型的复合路径比如“作者-论文-会议”或“用户-商品-品牌”。通过元路径可以把异构图转换成同构图两个作者如果通过论文发生过关联就建立一条同质边。之后再套 GCN 或 GAT 就行。元路径的难点在选择上。不同任务需要不同路径选错路径信息就丢了。我一般会先做数据探索统计各类关系的重要性再决定要用哪些元路径。5.5 动态图和异构图的组合场景真实项目里动态和异构经常同时存在。比如推荐系统节点类型多交互时间也在变知识图谱实体和关系类型多事实也在不断累积。处理组合场景没有统一银弹常见做法是先分层外层处理时间内层处理类型。先用时间编码编码节点状态再用异构图聚合模块处理多种边关系最后在输出层融合。也有框架直接支持这种场景比如 PyG 中的 temporal 相关模块和 DGL 的异构图 API但工程上仍需要自己设计采样和训练流程。6. 实操落地工具选型、最小代码和验证方法6.1 PyG 还是 DGL现在做 GNN 实践主流工具是 PyTorch Geometric 和 Deep Graph Library。PyG 和 PyTorch 结合紧密API 直观数据结构友好写完模型再写一个训练循环就能跑非常适合入门。DGL 在大规模图采样和异构图某些场景下优化更细能力也很强但学习成本略高。个人建议刚入门用 PyG少折腾工具细节把精力放在理解模型和数据上。等需要大规模分布式训练或特殊图结构时再评估 DGL。环境准备方面Python 3.8 以上先装 PyTorch版本要和你机器的 CUDA 对应再安装 torch-geometric 相关依赖。小数据集用 CPU 可以试但稍微大一点的图比如十几万节点强烈建议用 GPU否则训练速度会让人失去耐心。6.2 用 PyG 写一个两层 GCN 分类模型下面是一个最小可跑的 PyG 模型示例结构就是前面说的两层 GCNimport torch import torch.nn.functional as F from torch_geometric.nn import GCNConv class GCNNet(torch.nn.Module): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.conv1 GCNConv(in_dim, hidden_dim) self.conv2 GCNConv(hidden_dim, out_dim) def forward(self, x, edge_index): x self.conv1(x, edge_index) x F.relu(x) x F.dropout(x, trainingself.training) x self.conv2(x, edge_index) return xedge_index 是 [2, E] 的 LongTensor每一列是一条边的起点和终点。注意这个顺序第 0 行是源节点第 1 行是目标节点。如果你从原始邻接矩阵转换要确认方向、自环和去重都处理正确。6.3 训练循环、数据集划分和损失计算模型定义好之后训练循环和普通 PyTorch 差别不大model GCNNet(in_dimdataset.num_node_features, hidden_dim16, out_dimdataset.num_classes) optimizer torch.optim.Adam(model.parameters(), lr0.01) for epoch in range(200): model.train() optimizer.zero_grad() out model(data.x, data.edge_index) loss F.cross_entropy(out[data.train_mask], data.y[data.train_mask]) loss.backward() optimizer.step()这里只对 train_mask 选中的节点计算损失这是图半监督学习的常见写法整个图在训练时是可见的但损失和计算梯度只发生在有标签的节点上。测试时用 val_mask 和 test_mask 评估。数据划分要特别注意。节点分类里随机划分节点是常见做法但时序场景不行。如果数据本身有强烈时间依赖随机划分会让模型“提前看到未来”评估分数虚高。时间序列任务必须按时间切分训练集、验证集和测试集。6.4 怎么判断模型真的“学到了”不能只看训练集损失降下来就结束。要看验证集表现也要看输出有没有问题。常用指标根据任务不同变化任务常用指标节点分类Accuracy、Macro F1、Precision / Recall链接预测AUC、MRR、HitsK图分类Accuracy、F1回归任务看 MSE还要检查输出的类别分布是否畸形。如果模型把所有节点都预测成同一类准确率可能看起来还行尤其类别不平衡时实际没有任何用处。这种时候优先看混淆矩阵。7. 调试、踩坑和工程经验7.1 图数据最容易被忽略的四个坑第一节点 ID 不连续。很多原始数据里节点 ID 有空洞直接转成矩阵会出错。要先做重映射确保节点编号从 0 到 N-1 连续。第二自环缺失。处理过的图数据经常没有把自环加进去GCN 效果会受影响。PyG 里可以直接用 torch_geometric.utils.add_self_loops 补充。第三特征矩阵和邻接矩阵对不齐。数据源多个文件时节点顺序不一致是常见问题。特征矩阵第 i 行必须对应节点 i边索引里的节点 ID 也必须同一套编码。第四边方向搞反。有向图里方向是语义的一部分在 PyG 中 edge_index[0] 是源edge_index[1] 是目标方向错了整个图就反了。7.2 训练不收敛时按什么顺序排查训练不收敛先不要怀疑模型结构按下面顺序排查。第一步看 loss 是否是 NaN。如果是检查特征里有没有 NaN 或 inf学习率是不是太大输入特征要不要归一化。第二步看训练集准确率。如果训练集都不涨先检查标签取值是否在 0 到类别数减 1 之间模型输出维度是否等于类别数数据有没有喂错。第三步看验证集。训练集涨、验证集不涨是过拟合加 dropout、weight decay或者减少层数。第四步看 loss 曲线。震荡严重就降低学习率或者加学习率预热。最后才考虑换模型结构。用简单模型跑通一个流程再上复杂的模型这是最省时间的调试路径。7.3 过平滑、过拟合和欠拟合的处理过平滑是 GNN 特有的问题。层数超过四五层后所有节点表示趋向一致分类效果反而变差。缓解办法减层数加残差连接或者用 JK-Net 这种在每层输出之间做跳跃连接的方案。过拟合在 GNN 上也很常见尤其是小图。缓解措施dropout、weight decay、早停、减少模型层数和隐藏维度、增大训练节点比例。不要一上来就追求大模型。欠拟合时先确认特征是否有效。有个现象经常被忽略如果节点特征本身不含信号GNN 不一定比普通多层感知机 MLP 强。所以我会先跑一个 MLP baseline只看特征不看图结构再跑 GCN 对比。如果 GCN 没有明显提升说明图结构对当前任务帮助有限问题可能出在特征或任务定义上。7.4 真实项目落地前先想清楚这几件事第一先定任务和指标再定模型。很多人上来先选一个最新论文里的模型结果连 baseline 指标都没有定义好最后没法判断效果提升。第二把数据版本、特征版本、参数版本都记录清楚。图数据实验里一个负采样策略不同AUC 能差出好几个点。不记录版本问题出现时完全没法回溯。第三链接预测要统一负采样。预测正样本的同时要生成负样本负样本从哪里来、采样比例多少、是否做难负样本都会影响结果。第四注意推理阶段的工程限制。模型在离线评测里效果好线上可能因为新用户没有足够邻居导致表示不稳定。动态图场景里还有时间戳对齐问题训练时和推理时的时间分布不一致也会影响效果。第五如果需要上线提前设计好邻居采样和缓存策略。实时推理时每来一个新事件都要及时更新相关节点表示而不是每次都重算整个大图。这里的稳定性往往比模型精度更重要。GNN 不是万能工具但只要任务里确实存在图结构信息它通常能比普通模型更有效地利用这些关系。学习路径也简单先把消息传递的概念吃透再跑通 GCN然后用 GAT 做对比最后根据自己数据的特点去研究动态图、异构图和采样策略。踩过几次坑之后你会发现真正难的不是模型定义而是数据对不对、切分对不对、指标对不对。这三件事想清楚了大部分 GNN 项目都能稳定推进。
返回列表