引言
表格数据里样本彼此独立,图像数据里像素按网格排列,而现实世界的大量数据是图:社交网络、分子结构、知识图谱、推荐系统的用户-商品二分图。图没有固定的节点顺序,邻居数量也各不相同——卷积和循环网络都用不上。
**图神经网络(GNN)**给出了一套统一答案:让每个节点反复「收集邻居的信息」,从而学到融合了局部结构的表示。本文从图的基本表示讲起,拆解消息传递范式,再落到 GCN 与 GAT 两个经典模型,最后用 PyG 跑通节点分类。
前置:神经网络与反向传播基础见 https://plumephp.com/ml-neural-networks-basics/;嵌入向量的思想可对比 https://plumephp.com/ml-nlp-basics/ 中的词向量;推荐系统的图结构应用见 https://plumephp.com/ml-recommender-systems/。
目录
- 1. 图数据与图学习问题
- 2. 图的表示与邻接矩阵
- 3. 消息传递机制
- 4. GCN:从谱图卷积到一阶近似
- 5. GAT:注意力加权聚合
- 6. 节点分类与图分类
- 7. 用 PyG 跑通节点分类
- 8. 常见坑与调参清单
- 9. 总结
- 延伸阅读
1. 图数据与图学习问题
1.1 图的基本构成
节点是实体(用户、商品、原子),边是关系(关注、购买、化学键),节点与边都可以带属性向量。
1.2 三类典型任务
| 任务 | 预测对象 | 例子 |
|---|---|---|
| 节点级 | 每个节点的标签 | 用户是否流失、论文分类 |
| 图级 | 整张图的标签 | 分子是否有毒性 |
1.3 为什么不能直接用 MLP
把节点特征直接喂给 MLP,会丢掉结构信息:两个特征几乎一样的用户,可能因为朋友圈完全不同而有不同标签。GNN 的价值就是把「谁和谁相连」编码进表示。
2. 图的表示与邻接矩阵
2.1 邻接矩阵
N 个节点的图用 N×N 的矩阵 A 表示,A[i][j] = 1 表示 i 与 j 有边:
import numpy as np
# 3 个节点:0-1 相连,1-2 相连
A = np.array([[0, 1, 0], [1, 0, 1], [0, 1, 0]], dtype=float)
print(A.sum(axis=1)) # 每个节点的度:[1, 2, 1]
2.2 稀疏格式
真实图的邻接矩阵极其稀疏(百万节点、平均度几十)。存储用 COO 三元组:
import torch
# edge_index 形状 (2, E),第 0 行源、第 1 行目标(无向图两方向都存)
edge_index = torch.tensor([[0, 1, 1, 2],
[1, 0, 2, 1]], dtype=torch.long)
这是 PyG 的标准接口,比稠密矩阵省几个数量级的内存。
2.3 度矩阵与归一化
度矩阵 D 是对角矩阵,D[i][i] 等于节点 i 的度。归一化邻接矩阵是 GNN 的核心构件:
D^-1 A 行归一化,每行和为 1
D^-1/2 A D^-1/2 对称归一化,GCN 用的就是这个
对称归一化让高度节点不会因为邻居多而数值爆炸。
3. 消息传递机制
3.1 通用范式
几乎所有 GNN 都能写成三步循环(MPNN 框架):
对每一层 k:
1. 消息生成 m_ij = M(h_i, h_j, e_ij) 邻居 j 发来的消息
2. 消息聚合 m_i = AGG({m_ij : j 属于 N(i)}) 求和/均值/最大/注意力
3. 节点更新 h_i' = U(h_i, m_i) 通常接一个线性层+激活
3.2 直觉理解
一层消息传递 = 每个节点看一眼自己的直接邻居。堆两层,节点就能看到「邻居的邻居」,感受野为 2 跳;堆 k 层感受野为 k 跳。
3.3 手写一层消息传递
import torch
def simple_message_passing(h, edge_index):
"""最朴素的均值聚合:h_new[i] = mean(h[j] for j in N(i))"""
src, dst = edge_index[0], edge_index[1]
out = torch.zeros_like(h)
out.index_add_(0, dst, h[src]) # 把邻居特征累加到目标节点
deg = torch.zeros(h.size(0), device=h.device)
deg.index_add_(0, dst, torch.ones(src.size(0), device=h.device))
return out / deg.clamp(min=1).unsqueeze(-1)
index_add_ 是稀疏聚合的高效实现,等价于「按目标节点分组求和」。
3.4 自环:别忘了自己
只聚合邻居会丢失节点自身的信息。标准做法是给每个节点加一条指向自己的边:
def add_self_loops(edge_index, num_nodes):
loop = torch.arange(num_nodes, device=edge_index.device)
return torch.cat([edge_index, loop.unsqueeze(0).repeat(2, 1)], dim=1)
4. GCN:从谱图卷积到一阶近似
4.1 公式
GCN 的逐层传播规则只有一行:
H' = sigma( A_hat H W )
A_hat = D_hat^-1/2 (A + I) D_hat^-1/2 ← 加自环后对称归一化
其中 H 是节点特征矩阵,W 是可学习权重,sigma 是激活函数。
4.2 这个公式从哪来
理论上它源自谱图卷积的切比雪夫多项式近似,取一阶截断并做重归一化技巧(renormalization trick)得到。工程上不必深究推导,记住三点即可:
- 加自环:聚合时包含自身;
- 对称归一化:按度数平衡邻居贡献;
- 线性变换 + 激活:与普通神经网络层一致。
4.3 用 PyG 实现一个 GCN 层
import torch.nn as nn
from torch_geometric.nn import GCNConv
class GCN(nn.Module):
def __init__(self, in_dim, hidden, out_dim):
super().__init__()
self.conv1 = GCNConv(in_dim, hidden)
self.conv2 = GCNConv(hidden, out_dim)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index).relu()
x = nn.functional.dropout(x, p=0.5, training=self.training)
return self.conv2(x, edge_index)
4.4 层数不是越多越好
| 层数 | 感受野 | 效果 |
|---|---|---|
| 1 | 1 跳 | 欠拟合,结构信息不足 |
| 2 | 2 跳 | 多数任务的甜点 |
| 4+ | 4 跳以上 | 过平滑,节点表示趋同 |
**过平滑(over-smoothing)**是 GNN 特有的病:层数一多,所有节点的表示都被邻居平均得趋于一致,分类能力崩塌。实践中 2~3 层最常见。
5. GAT:注意力加权聚合
5.1 均值聚合的问题
GCN 对邻居一视同仁(只按度数缩放)。但现实中「最重要的那个邻居」往往才是关键——引用网络里,被大牛引用的论文权重应该更高。
5.2 注意力系数
GAT 为每条边学一个注意力权重:
e_ij = LeakyReLU( a^T [ W h_i || W h_j ] ) 拼接后过单层网络
alpha_ij = softmax_j(e_ij) 在同一节点的邻居间归一化
h_i' = sigma( sum_j alpha_ij W h_j ) 加权求和
关键细节:softmax 是在每个节点的邻居集合上做的,不是全图。
5.3 多头注意力
和 Transformer 一样,GAT 用多头增强稳定性:中间层多头结果拼接(concat),输出层多头结果平均(mean)。
from torch_geometric.nn import GATConv
class GAT(nn.Module):
def __init__(self, in_dim, hidden, out_dim, heads=8):
super().__init__()
self.conv1 = GATConv(in_dim, hidden, heads=heads, dropout=0.6)
# 输出层用单头并平均,避免维度爆炸
self.conv2 = GATConv(hidden * heads, out_dim, heads=1,
concat=False, dropout=0.6)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index).elu()
return self.conv2(x, edge_index)
5.4 GCN 与 GAT 对比
| 维度 | GCN | GAT |
|---|---|---|
| 聚合方式 | 度归一化加权 | 学习出的注意力 |
| 参数量 | 少 | 多(注意力参数) |
| 表达能力 | 中 | 强 |
| 适用 | 同质图基线 | 邻居重要性差异大 |
6. 节点分类与图分类
6.1 节点分类:半监督为主
节点分类的典型设定是半监督:只标注少量节点,靠结构把标签传播开。损失只在小部分有标签的节点上计算:
criterion = nn.CrossEntropyLoss()
def train_step(model, data, optimizer):
model.train(); optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = criterion(out[data.train_mask], data.y[data.train_mask])
loss.backward(); optimizer.step()
return loss.item()
注意 out[data.train_mask]:没有标签的节点不参与损失,但它们的信息通过邻接传播进了有标签节点的表示里,这就是半监督的魔法。
6.2 图分类:需要全局池化
整张图输出一个标签,必须把节点表示聚合成图表示:
全局池化:h_G = mean / sum / max({h_i})
层次池化:DiffPool / TopKPool,边池化边学结构
from torch_geometric.nn import global_mean_pool
class GraphClassifier(nn.Module):
def __init__(self, in_dim, hidden, num_classes):
super().__init__()
self.conv1 = GCNConv(in_dim, hidden)
self.conv2 = GCNConv(hidden, hidden)
self.head = nn.Linear(hidden, num_classes)
def forward(self, x, edge_index, batch):
x = self.conv1(x, edge_index).relu()
x = self.conv2(x, edge_index).relu()
x = global_mean_pool(x, batch) # batch 指明每个节点属于哪张图
return self.head(x)
6.3 两类任务对比
| 维度 | 节点分类 | 图分类 |
|---|---|---|
| 输出 | 每节点一个标签 | 每图一个标签 |
| 池化 | 不需要 | 必需 |
| 代表数据集 | Cora、PubMed | MUTAG、PROTEINS |
| 常见任务 | 用户画像、论文分类 | 分子性质、代码分类 |
7. 用 PyG 跑通节点分类
7.1 加载数据
import torch
from torch_geometric.datasets import Planetoid
from torch_geometric.transforms import NormalizeFeatures
dataset = Planetoid(root="/tmp/Cora", name="Cora",
transform=NormalizeFeatures())
data = dataset[0]
# Data(x=[2708, 1433], edge_index=[2, 10556], y=[2708],
# train_mask=[2708], val_mask=[2708], test_mask=[2708])
Cora 是 2708 篇论文、1433 维词袋特征、7 个类别、10556 条引用边。
7.2 训练与评估
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = GCN(dataset.num_features, 16, dataset.num_classes).to(device)
data = data.to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
def evaluate(mask):
model.eval()
with torch.no_grad():
pred = model(data.x, data.edge_index).argmax(dim=1)
return (pred[mask] == data.y[mask]).float().mean().item()
best_val = 0
for epoch in range(200):
loss = train_step(model, data, optimizer)
val_acc = evaluate(data.val_mask)
if val_acc > best_val:
best_val = val_acc
torch.save(model.state_dict(), "best.pt")
if epoch % 20 == 0:
print(f"epoch {epoch:3d} loss {loss:.4f} val {val_acc:.4f}")
model.load_state_dict(torch.load("best.pt"))
print("test acc:", round(evaluate(data.test_mask), 4))
# 两层 GCN 在 Cora 上通常能到 0.80 左右
7.3 结果解读
GCN(2 层) Cora 测试准确率 约 0.80
GAT(2 层 8 头) 约 0.82~0.83
MLP(忽略结构) 约 0.55~0.60
MLP 与 GCN 的差距就是结构信息的价值——同样的特征,加了邻接关系后准确率提升 20 多个百分点。
8. 常见坑与调参清单
| 现象 | 根因 | 处理 |
|---|---|---|
| 训练准确率高、测试低 | 过拟合(层多、参数多) | 减层、加 dropout、加 weight_decay |
| 层数增加反而变差 | 过平滑 | 回到 2~3 层,或加残差连接 |
| loss 不下降 | 忘记加自环 / 特征未归一化 | 检查 NormalizeFeatures 与 conv 实现 |
| 边方向搞反 | 有向图 source 与 target 弄混 | 明确「谁聚合谁」的语义 |
| 大图 OOM | 全图训练放不下 | 邻居采样(GraphSAGE 的 NeighborLoader) |
| 度分布极端 | 少数超级节点主导 | 用对称归一化,或对度做截断 |
8.1 超参经验值
hidden_dim 16~256(Cora 上 16 就够) layers 2~3 dropout 0.5~0.6
lr 0.005~0.01 weight_decay 5e-4 epochs 200(早停看验证集)
8.2 大图怎么办
全图训练要求整张图和全部特征驻留显存,百万节点就不行了。三种方案:
- 邻居采样:每批只采 K 跳邻居子图(GraphSAGE);
- 图聚类:先用 METIS 切成子图,分批训练(Cluster-GCN);
- 特征降维:1433 维词袋压到 128 维,显存立省 10 倍。
9. 总结
9.1 学习路径
图的表示 → 消息传递范式(消息-聚合-更新)
→ GCN(归一化 + 自环,2 层甜点)→ GAT(注意力加权)
→ 节点分类(半监督掩码)/ 图分类(全局池化)
→ 大图采样(NeighborLoader)
9.2 关键决策点
| 问题 | 选择 |
|---|---|
| 邻居重要性差异大 | GAT |
| 只想快速跑通基线 | GCN 2 层 |
| 层数加到 4 层掉点 | 过平滑,退回 2~3 层或加残差 |
| 图有数百万节点 | 邻居采样或 Cluster-GCN |
| 图级任务 | 必须加全局池化层 |
| 节点无特征 | 用度、PageRank 等结构特征兜底 |
9.3 一句话心法
GNN 的全部秘密就是「反复聚合邻居」——理解了三步消息传递,GCN、GAT、GraphSAGE 都只是聚合函数的不同选择而已。
延伸阅读
- https://plumephp.com/ml-neural-networks-basics/ — 神经网络与反向传播基础
- https://plumephp.com/ml-nlp-basics/ — 词嵌入与表示学习思想
- https://plumephp.com/ml-recommender-systems/ — 二分图上的推荐与召回排序
- https://plumephp.com/ml-unsupervised-clustering/ — 图聚类与社区发现的对照
- PyTorch Geometric 官方文档
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。