图数据的价值藏在关系里:用户与商品的购买边、分子里原子间的化学键、知识图谱中的实体关系。卷积网络只能处理规则网格,循环网络只能处理序列,一旦数据变成任意拓扑的图,就需要一套新的计算范式——图神经网络。
从图结构到图神经网络
图由顶点集合与边集合构成,可以带方向、带权重、带类型。真正让图学习困难的是排列不变性:同一个图,邻接矩阵换一种节点编号就完全变了,但图的语义没变。模型必须对节点顺序不敏感。
图数据的三要素
- 节点特征矩阵
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 层的更新分三步:
- 消息生成:每条边
(u, v)生成一条消息m_{u→v} = M(h_u^{k-1}, h_v^{k-1}, e_{uv})。 - 消息聚合:把邻居发来的消息汇总
a_v = AGG({m_{u→v} : u ∈ N(v)})。 - 状态更新:
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 的对比
| 维度 | GCN | GAT |
|---|---|---|
| 邻居权重 | 由度数固定决定 | 由注意力学习 |
| 参数量 | 少 | 多一个注意力向量 |
| 显存 | 低 | 多头导致显存翻倍 |
| 小图表现 | 好 | 好 |
| 大图表现 | 稳 | 注意力开销大,需采样 |
| 可解释性 | 弱 | 注意力权重可解释 |
实践中:图规模小、要可解释性,选 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] | 逐层递减 |
| dropout | 0.5 | 图数据极易过拟合 |
| 学习率 | 0.01 | Adam,配 weight decay |
| 归一化 | BatchNorm | 每层后加 |
另外两个高频问题:验证集指标远高于测试集往往是图划分有偏(如按社区划分导致分布不一致);loss 不降先检查邻接矩阵方向——edge_index 的 [0] 行是源节点还是目标节点搞反,模型会学成反向传播。
总结
图神经网络的主线是「用不变聚合函数,把邻居信息逐层汇聚」:消息传递给出统一范式,GCN 用对称归一化做加权平均,GraphSAGE 用采样与聚合函数实现归纳学习,GAT 用注意力学边权。深度不是解药,2~4 层配残差与 JKNet 已足够;规模才是真问题,邻居采样、GraphSAINT 与 Cluster-GCN 是把图塞进显存的三种武器。落地时,先保证训练与推理的采样一致、切分无泄漏,再谈调参——图模型的绝大部分线上事故,都出在数据管道而不是模型结构。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。