图神经网络实战:消息传递、GCN/GAT/GraphSAGE 与大图训练

图神经网络把深度学习从规则网格推广到任意拓扑结构,本文系统讲解消息传递范式的三个阶段与数学形式、GCN 的对称归一化传播推导、GraphSAGE 的邻居采样与归纳学习、GAT 的注意力加权聚合、过平滑的成因与残差跳跃缓解、大图训练的邻居采样与 GraphSAINT/Cluster-GCN、节点边图三类任务与图级读出、与 Neo4j 等图数据库的特征协同,以及生产落地中的数据泄漏与邻居爆炸踩坑。

图数据的价值藏在关系里:用户与商品的购买边、分子里原子间的化学键、知识图谱中的实体关系。卷积网络只能处理规则网格,循环网络只能处理序列,一旦数据变成任意拓扑的图,就需要一套新的计算范式——图神经网络。

从图结构到图神经网络

图由顶点集合与边集合构成,可以带方向、带权重、带类型。真正让图学习困难的是排列不变性:同一个图,邻接矩阵换一种节点编号就完全变了,但图的语义没变。模型必须对节点顺序不敏感。

图数据的三要素

  • 节点特征矩阵 X ∈ R^{N×F}:每个节点的属性,如用户年龄、物品类目。
  • 邻接矩阵 A ∈ {0,1}^{N×N}:谁与谁相连,可带权重。
  • 边特征 E:可选,如交易金额、时间戳。

工业界图往往是异质图:节点有多种类型,边有多种语义。异质图需要按类型分别定义聚合函数,常见做法是先转成同质子图,或使用 HGT 之类的异质模型。

为什么卷积与循环网络不适用

  • CNN 依赖平移不变性:图没有固定的邻域顺序,3×3 卷积核无处安放。
  • RNN 依赖线性序列:图有环、有分支,且节点没有天然的先后顺序。
  • MLP 只看节点自身特征:完全忽略了结构信息,等价于丢掉了图的一半价值。

三种主流流派

流派核心思想代表方法
谱域方法在拉普拉斯特征空间做卷积ChebNet、GCN
空域方法直接聚合邻居消息GraphSAGE、GAT、MPNN
随机游走用游走序列学嵌入DeepWalk、node2vec

当前主流是空域方法,因为它天然支持归纳学习、易于采样、工程实现简单。谱域方法提供了理论直觉,GCN 正是从谱域简化而来的空域实现。

消息传递范式

Gilmer 等人在 2017 年提出的 Message Passing Neural Network(MPNN) 统一了几乎所有 GNN:无论模型叫什么名字,本质都是「邻居发消息、节点收消息、更新自己的状态」。

消息传递的三个阶段

对一个节点 v,第 k 层的更新分三步:

  1. 消息生成:每条边 (u, v) 生成一条消息 m_{u→v} = M(h_u^{k-1}, h_v^{k-1}, e_{uv})。
  2. 消息聚合:把邻居发来的消息汇总 a_v = AGG({m_{u→v} : u ∈ N(v)})。
  3. 状态更新:h_v^k = U(h_v^{k-1}, a_v)。

关键约束是聚合函数必须对邻居排列不变:求和、求均值、取最大都满足,拼接不满足。

数学形式

以最通用的形式写出:

h_v^{(k)} = σ( W_self · h_v^{(k-1)} + W_neigh · AGG_{u∈N(v)} h_u^{(k-1)} )

其中 W_self 与 W_neigh 是可学习参数,σ 是激活函数。这个形式是 GCN、GraphSAGE、GIN 的共同骨架,差异只在 AGG 与是否归一化。

PyTorch 最小实现

用 PyG 的 MessagePassing 基类可以几十行写出一个自定义层:

import torch
from torch import nn
from torch_geometric.nn import MessagePassing
from torch_geometric.utils import add_self_loops

class SimpleMP(MessagePassing):
    def __init__(self, in_dim, out_dim):
        super().__init__(aggr="add")          # 聚合方式:add / mean / max
        self.lin_self = nn.Linear(in_dim, out_dim)
        self.lin_neigh = nn.Linear(in_dim, out_dim)

    def forward(self, x, edge_index):
        edge_index, _ = add_self_loops(edge_index, num_nodes=x.size(0))
        x = self.lin_self(x) + self.lin_neigh(x)
        return self.propagate(edge_index, x=x)   # 触发 message/aggregate/update

    def message(self, x_j):                      # x_j 是邻居特征
        return x_j

x_j 是 PyG 的约定:带 _j 后缀表示「邻居端」特征,_i 表示「中心端」。理解这个约定,是读懂所有 PyG 源码的前提。

GCN:从谱域到空域的简化

GCN 的原始动机是谱图卷积:把图信号做图傅里叶变换,在频域做乘法再变换回来。但完整谱卷积需要特征分解,复杂度 O(N³),无法实用。

谱图卷积的直觉

Kipf 与 Welling 用一阶切比雪夫多项式近似谱卷积,得到极简的传播规则。直觉上,它等价于「每个节点取自己和邻居特征的加权平均」,权重由度归一化决定。

对称归一化的传播规则

GCN 的核心公式:

H^{(k)} = σ( D^{-1/2} Ã D^{-1/2} H^{(k-1)} W^{(k-1)} )

其中 Ã = A + I 是加上自环的邻接矩阵,D̃ 是 Ã 的度矩阵。自环保证节点在聚合时保留自身信息,D^{-1/2} 的对称归一化避免高度节点主导。

import torch.nn.functional as F
from torch_geometric.nn import GCNConv

class GCN(nn.Module):
    def __init__(self, in_dim, hid_dim, out_dim, dropout=0.5):
        super().__init__()
        self.conv1 = GCNConv(in_dim, hid_dim)
        self.conv2 = GCNConv(hid_dim, out_dim)
        self.dropout = dropout

    def forward(self, x, edge_index):
        x = F.relu(self.conv1(x, edge_index))
        x = F.dropout(x, p=self.dropout, training=self.training)
        return self.conv2(x, edge_index)

显存与复杂度

GCN 单层的计算量约为 O(E × F × F'),显存峰值来自稀疏矩阵乘。全图训练时,邻接矩阵的稀疏结构用 COO 存储,边数 E 决定了内存地板。Cora 数据集只有 5429 条边,而工业级图动辄上亿条边,这就是大图训练必须采样的原因。

GraphSAGE:采样与归纳学习

GCN 是直推式的:训练时见过哪些节点,就只能预测哪些节点。新增一个节点,需要重训。GraphSAGE 用「学习聚合函数」替代「学习每个节点的嵌入」,从而实现归纳式学习。

直推式与归纳式

  • 直推式(Transductive):为每个节点学一个嵌入向量,参数数量随节点数增长。新节点无嵌入,无法预测。
  • 归纳式(Inductive):学习的是「如何从邻居特征聚合」的规则,参数与节点数无关。新节点只要有特征和邻居,立刻可以推理。

工业场景几乎都需要归纳式——用户天天新增,不可能每天重训全图。

三种聚合器

GraphSAGE 论文对比了三种聚合函数:

聚合器形式特点
Mean邻居特征逐元素平均简单稳定,最常用
LSTM邻居序列过 LSTM 再取输出表达强但不对称,需随机打乱
Pooling邻居过 MLP 后逐维取 max表达强,计算略重

聚合后与自身特征拼接,再过一个线性层:

h_v^k = σ( W · CONCAT( h_v^{k-1}, AGG_{u∈N(v)} h_u^{k-1} ) )

邻居采样

GraphSAGE 的关键工程创新是固定扇出采样:每层只随机采 S_k 个邻居。若扇出为 [10, 25],则两层的计算树大小上限为 10 × 25 = 250 个节点,与图的度数无关。这把「随度数爆炸」变成「常数规模」,是大图训练能够 mini-batch 化的基石。

from torch_geometric.nn import SAGEConv
from torch_geometric.loader import NeighborLoader

conv = SAGEConv(in_dim, hid_dim, aggr="mean")

loader = NeighborLoader(
    data,
    num_neighbors=[10, 25],      # 每层采样扇出
    batch_size=1024,
    input_nodes=data.train_mask,
    shuffle=True,
)

num_neighbors 是 GraphSAGE 最重要的超参:越大越接近全图、方差越小,但显存与耗时越高。

GAT:注意力加权的邻域聚合

GCN 与 GraphSAGE 对邻居一视同仁(或仅按度数加权)。现实中邻居的重要性天差地别——交易网络里,一笔大额转账的邻居比一笔小额更有信息量。GAT(Graph Attention Network) 用注意力机制自动学出每条边的权重。

注意力系数计算

对边 (i, j),先算未归一化的注意力得分:

e_{ij} = LeakyReLU( a^T · [ W h_i || W h_j ] )

再用 softmax 在邻居维度归一化:

α_{ij} = softmax_j(e_{ij}) = exp(e_{ij}) / Σ_{k∈N(i)} exp(e_{ik})

最终聚合为 h_i' = σ( Σ_{j∈N(i)} α_{ij} W h_j )。

多头注意力

为稳定训练,GAT 用多头注意力:K 个独立的注意力头并行计算,中间层拼接、输出层取平均。这与 Transformer 的多头设计完全同源。

from torch_geometric.nn import GATConv

class GAT(nn.Module):
    def __init__(self, in_dim, hid_dim, out_dim, heads=8):
        super().__init__()
        self.conv1 = GATConv(in_dim, hid_dim, heads=heads, dropout=0.6)
        # 输出层单头,避免维度爆炸
        self.conv2 = GATConv(hid_dim * heads, out_dim, heads=1,
                             concat=False, dropout=0.6)

    def forward(self, x, edge_index):
        x = F.elu(self.conv1(x, edge_index))
        return self.conv2(x, edge_index)

与 GCN 的对比

维度GCNGAT
邻居权重由度数固定决定由注意力学习
参数量少多一个注意力向量
显存低多头导致显存翻倍
小图表现好好
大图表现稳注意力开销大,需采样
可解释性弱注意力权重可解释

实践中:图规模小、要可解释性,选 GAT;图规模大、要稳定吞吐,选 GraphSAGE。GCN 则是两者的折中基线,永远值得先跑一遍。

过平滑与深度受限

GNN 最反直觉的一点是:层数不是越深越好。图像里 100 层 ResNet 很常见,图里 4 层往往就饱和,8 层以上开始掉点。

过平滑现象

每一层聚合都在做邻域平均,等价于一次低通滤波。层数一多,所有节点的表示被反复平滑,最终收敛到同一个向量——所有节点变得无法区分,分类性能崩盘。这就是过平滑(Over-smoothing)。

一个直观的度量是节点表示的两两余弦相似度:层数增加时它会持续上升,逼近 1 时模型已失效。

残差与跳跃连接

缓解过平滑的主流手段:

  • 残差连接:h^k = h^{k-1} + Δh,保留原始信息。
  • JKNet(Jumping Knowledge):把每一层的输出都收集起来,最后拼接或取最大,让模型自己选择用几层。
  • DropEdge:训练时随机丢弃一部分边,减缓平滑速度,兼作正则化。
  • PairNorm:每层后做归一化,把节点表示的「总距离」拉回常数。
class JKNet(nn.Module):
    def __init__(self, in_dim, hid_dim, out_dim, num_layers=4):
        super().__init__()
        self.convs = nn.ModuleList(
            [GCNConv(in_dim if i == 0 else hid_dim, hid_dim)
             for i in range(num_layers)]
        )
        self.jk = nn.Linear(hid_dim * num_layers, out_dim)

    def forward(self, x, edge_index):
        hs = []
        for conv in self.convs:
            x = F.relu(conv(x, edge_index))
            hs.append(x)
        return self.jk(torch.cat(hs, dim=-1))

常用缓解手段

工程经验:GNN 的有效深度通常是 2~4 层。想让模型看到更远的邻居,靠的是「加层」以外的办法——比如在图上预计算多跳邻接、用图数据库先做子图抽取、或改用能传递长程信息的图 Transformer。

大图训练与邻居采样

全图训练需要把整个图放进显存。当边数上亿时,这条路走不通,必须把图切分成 mini-batch。

全图训练与 mini-batch

  • 全图训练:一次前向用整张图,梯度最准,但显存 O(N+E)。适合百万节点以内的图。
  • 邻居采样:每个 batch 抽一批目标节点,再按扇出采邻居,显存与图规模解耦。适合工业级大图。

GraphSAINT 与 Cluster-GCN

除了逐节点采样,还有两类更高效的子图采样方法:

方法采样粒度特点
NeighborLoader节点计算树通用,PyG 默认
GraphSAINT边/节点/随机游走子图子图内完整,方差小
Cluster-GCN图聚类簇簇内边密集,通信高效

Cluster-GCN 先用 METIS 把图聚成若干簇,每个 batch 取几个簇。因为簇内边远多于跨簇边,避免了邻居采样那种「采出来的子图边很稀疏」的浪费,训练速度快很多。

from torch_geometric.loader import ClusterData, ClusterLoader

cluster_data = ClusterData(data, num_parts=1500, recursive=False)
loader = ClusterLoader(cluster_data, batch_size=20, shuffle=True)

for batch in loader:
    out = model(batch.x, batch.edge_index)
    loss = F.cross_entropy(out[batch.train_mask], batch.y[batch.train_mask])

采样带来的方差

采样是有偏估计的方差来源:扇出越小,采样的邻居子集波动越大,梯度噪声越重。实践中的补偿手段:

  • 扇出逐层递减:第一层采多(如 15),越靠近目标节点采越多,远端可以少采。
  • 增大 batch:用更大的 batch 平均掉采样噪声。
  • 重要性采样:按度数或边权设计采样概率,降低估计偏差。

节点边图三类任务

GNN 的下游任务按预测粒度分三类,损失函数与读出方式各不相同。

节点级任务

给定节点表示 h_v,接一个分类或回归头。典型场景:用户画像分类、论文主题分类、欺诈节点识别。

logits = model(x, edge_index)
loss = F.cross_entropy(logits[train_mask], y[train_mask])

注意 mask 的划分:训练/验证/测试集必须按节点划分,且要检查是否存在「训练节点的邻居大量落在测试集」导致的泄漏。

边级与链接预测

边级任务预测一条边的属性(如交易是否欺诈)。做法是把两端节点表示拼接或做内积,再过一个 MLP:

z = model(x, edge_index)
edge_emb = torch.cat([z[edge_label_index[0]], z[edge_label_index[1]]], dim=-1)
logits = edge_classifier(edge_emb)

链接预测更常见:预测两个节点之间是否存在边。它需要负采样构造负边,并用 AUC 或 Hits@K 评估。关键细节是负采样分布——均匀负采样会让任务过于简单,工业界常用按度数采样的负例。

图级任务与读出

图级任务(如分子性质预测、图分类)需要一个读出(Readout) 函数把节点表示汇总成图表示:

h_G = READOUT({ h_v : v ∈ V })

常用读出是全局求和、全局平均或全局最大,也可以引入层次化池化(DiffPool、TopK Pool)学习一个可微的软聚类。

from torch_geometric.nn import global_mean_pool

class GraphClassifier(nn.Module):
    def __init__(self, in_dim, hid_dim, num_classes):
        super().__init__()
        self.conv1 = GCNConv(in_dim, hid_dim)
        self.conv2 = GCNConv(hid_dim, hid_dim)
        self.head = nn.Linear(hid_dim, num_classes)

    def forward(self, x, edge_index, batch):
        x = F.relu(self.conv1(x, edge_index))
        x = F.relu(self.conv2(x, edge_index))
        g = global_mean_pool(x, batch)     # 按 batch 向量做分段平均
        return self.head(g)

图级任务尤其要注意 GIN 的教训:求和聚合比平均聚合表达力更强,因为平均会丢失节点数量信息。若任务对图的规模敏感(如分子大小影响性质),用求和而非平均。

与图数据库和特征工程的协同

GNN 的输入不是原始数据库,而是从图数据库抽取、加工后的张量。这一环决定了模型能否真正上线。

从 Neo4j 导出子图

工业图通常存在 Neo4j 之类的图数据库中。训练时导出子图,推理时按需拉取邻居:

# 用 Cypher 抽取目标节点及其两跳邻居,导出为边表
# MATCH (u:User)-[r:RATED]->(i:Item)
# WHERE u.id IN $seed_ids
# RETURN u.id AS src, i.id AS dst, r.score AS weight

导出的边表与节点特征表,在 Python 侧组装成 PyG 的 Data 对象。大图导出要分批拉取,避免一次性把内存打满。

特征工程与 ID 嵌入

节点特征通常有三类来源:

  • 属性特征:用户年龄、物品价格,直接归一化。
  • 统计特征:节点度数、PageRank、社区编号,用图算法预计算。
  • ID 嵌入:为高价值节点学一个可训练嵌入,但要注意冷启动——新节点没有历史 ID。

一个常见的融合方式是属性特征 + 预计算的图结构特征 + 可训练 ID 嵌入三者拼接。纯 ID 嵌入在冷启动场景会失效,必须保留属性通路。

在线推理的邻居获取

在线推理时,每个请求需要实时拉取目标节点的 k 跳邻居。这里的工程约束很硬:

  • 延迟预算:端到端 50ms 内,留给图查询的可能只有 10ms。
  • 邻居缓存:热点节点的邻居子图常驻 Redis,命中率高。
  • 扇出截断:在线场景只能采少量邻居(如 5~10),比训练时的扇出小得多。

训练与推理的采样分布不一致(train-serving skew)是线上掉点最常见的根因,务必在离线用线上同款采样逻辑复现一次指标。

生产实践与踩坑

GNN 从 notebook 到线上,坑比模型本身多。

数据泄漏

最隐蔽的坑。三类典型泄漏:

  • 时间泄漏:用未来的边预测过去的事件。必须按时间切分边,训练集只用 t < T 的边。
  • 特征泄漏:目标特征被间接编码进了节点特征。
  • 标签传播泄漏:训练时若把测试节点的边也放进图里做聚合,测试节点会「看到」自己的标签。

邻居爆炸

跳数增加时,邻居数指数增长。2 跳扇出 10 会触及 100 个节点,3 跳就是 1000。控制手段:

# 每层扇出显式声明,避免默认全采
num_neighbors=[15, 10, 5]   # 近端多采,远端少采
# 配合 pyg-lib 的采样内核,可显著降低 CPU 采样开销

常见调参清单

参数建议起点说明
层数2~3超过 4 层先怀疑过平滑
隐藏维度64~256与节点特征维度同量级
扇出[15, 10, 5]逐层递减
dropout0.5图数据极易过拟合
学习率0.01Adam,配 weight decay
归一化BatchNorm每层后加

另外两个高频问题:验证集指标远高于测试集往往是图划分有偏(如按社区划分导致分布不一致);loss 不降先检查邻接矩阵方向——edge_index 的 [0] 行是源节点还是目标节点搞反,模型会学成反向传播。

总结

图神经网络的主线是「用不变聚合函数,把邻居信息逐层汇聚」:消息传递给出统一范式,GCN 用对称归一化做加权平均,GraphSAGE 用采样与聚合函数实现归纳学习,GAT 用注意力学边权。深度不是解药,2~4 层配残差与 JKNet 已足够;规模才是真问题,邻居采样、GraphSAINT 与 Cluster-GCN 是把图塞进显存的三种武器。落地时,先保证训练与推理的采样一致、切分无泄漏,再谈调参——图模型的绝大部分线上事故,都出在数据管道而不是模型结构。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. GPU 共享与调度:MPS、MIG 与多租户隔离
  2. 异构推理硬件:ROCm、Intel 与国产 NPU 适配实践
  3. 前缀缓存与语义缓存:KV 复用与重复计算消除