图神经网络入门实战:消息传递、GCN 与 GAT

从零理解图神经网络:图的表示与邻接矩阵、消息传递机制的通用范式、GCN 的谱图卷积一阶近似与归一化、GAT 的注意力加权聚合、节点分类与图分类的任务差异,并用 PyTorch Geometric 在 Cora 数据集上跑通完整训练。

引言

表格数据里样本彼此独立,图像数据里像素按网格排列,而现实世界的大量数据是图:社交网络、分子结构、知识图谱、推荐系统的用户-商品二分图。图没有固定的节点顺序,邻居数量也各不相同——卷积和循环网络都用不上。

**图神经网络(GNN)**给出了一套统一答案:让每个节点反复「收集邻居的信息」,从而学到融合了局部结构的表示。本文从图的基本表示讲起,拆解消息传递范式,再落到 GCN 与 GAT 两个经典模型,最后用 PyG 跑通节点分类。

前置:神经网络与反向传播基础见 https://plumephp.com/ml-neural-networks-basics/;嵌入向量的思想可对比 https://plumephp.com/ml-nlp-basics/ 中的词向量;推荐系统的图结构应用见 https://plumephp.com/ml-recommender-systems/。


目录


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)得到。工程上不必深究推导,记住三点即可:

  1. 加自环:聚合时包含自身;
  2. 对称归一化:按度数平衡邻居贡献;
  3. 线性变换 + 激活:与普通神经网络层一致。

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 层数不是越多越好

层数感受野效果
11 跳欠拟合,结构信息不足
22 跳多数任务的甜点
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 对比

维度GCNGAT
聚合方式度归一化加权学习出的注意力
参数量少多(注意力参数)
表达能力中强
适用同质图基线邻居重要性差异大

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、PubMedMUTAG、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 大图怎么办

全图训练要求整张图和全部特征驻留显存,百万节点就不行了。三种方案:

  1. 邻居采样:每批只采 K 跳邻居子图(GraphSAGE);
  2. 图聚类:先用 METIS 切成子图,分批训练(Cluster-GCN);
  3. 特征降维: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 官方文档

继续阅读

探索更多技术文章

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

全部文章 返回首页

「ml」更多文章

  1. 特征平台与训练服务一致性:时间点正确性与特征回填
  2. 检索增强生成与向量检索实战:Embedding、HNSW 与重排
  3. 实验管理与可复现:MLflow、版本化与模型注册表