第17篇:图神经网络(GNN):处理非欧几里得数据的利器
·
摘要:
本文系统讲解图神经网络(GNN)的核心思想,深入解析图卷积网络(GCN)的谱域与空域原理、图注意力网络(GAT)的注意力机制、图自编码器(GAE)与图生成。结合PyTorch Geometric(PyG)库,实战Cora引文网络节点分类与链接预测。帮助学习者掌握处理社交网络、分子结构等图数据的“图”灵工具。
一、为什么需要GNN?
传统深度学习(CNN、RNN)处理欧几里得数据(如图像、序列),但许多现实数据是图结构(Graph):
- ✅ 社交网络:用户为节点,关注关系为边
- ✅ 分子结构:原子为节点,化学键为边
- ✅ 知识图谱:实体为节点,关系为边
- ✅ 推荐系统:用户-商品二部图
✅ 图数据具有不规则结构(变长邻居、无固定顺序),CNN/RNN难以直接应用。
二、图的基本概念
2.1 图的表示
- 图
G = (V, E)V:节点集合(|V| = N)E:边集合
- 邻接矩阵
A ∈ ℝ^{N×N}:Aᵢⱼ = 1表示节点i与j相连
- 节点特征矩阵
X ∈ ℝ^{N×F}:- 每行
Xᵢ是节点i的F维特征
- 每行
三、图神经网络核心思想:消息传递
GNN通过消息传递(Message Passing)更新节点表示:
- 聚合(Aggregate):收集邻居信息
- 更新(Update):结合自身状态更新表示
通用公式:
hᵥ^(k) = UPDATEₖ( hᵥ^(k-1), AGGREGATEₖ({hᵤ^(k-1) | u ∈ N(v)}) )
hᵥ^(k):节点v在第k层的嵌入N(v):v的邻居集合k:GNN层数
✅ 信息通过边在网络中“流动”。
四、图卷积网络(GCN)
4.1 空域解释(Spatial)
GCN的聚合函数为加权平均:
hᵥ^(k) = σ( Wₖ · mean({hᵤ^(k-1) | u ∈ N(v) ∪ {v}}) )
- 包含自环(
v自身) Wₖ:可训练权重σ:激活函数(如ReLU)
4.2 谱域解释(Spectral)
基于图信号处理:
- 图拉普拉斯矩阵
L = D - A(D为度矩阵) - GCN简化为:
H^(k) = σ( Â H^(k-1) Wₖ )Â = D̃⁻¹ᐟ² Ã D̃⁻¹ᐟ²:归一化邻接矩阵(带自环)Ã = A + ID̃:Ã的度矩阵
✅ GCN是谱图卷积的一阶近似。
五、图注意力网络(GAT)
5.1 核心思想
使用注意力机制为不同邻居分配不同权重。
5.2 计算步骤
-
计算注意力系数:
eᵢⱼ = a(W hᵢ, W hⱼ)a:注意力函数(如单层MLP + LeakyReLU)
-
归一化(softmax):
αᵢⱼ = softmaxⱼ(eᵢⱼ) = exp(eᵢⱼ) / Σ_{k∈N(i)} exp(eᵢₖ) -
加权求和:
hᵢ' = σ( Σ_{j∈N(i)} αᵢⱼ W hⱼ )
✅ GAT无需预先知道图结构,可处理动态图。
六、图自编码器(GAE)与图生成
6.1 图自编码器(GAE)
- 编码器:GNN → 学习节点嵌入
Z = GCN(X, A) - 解码器:从嵌入重构邻接矩阵:
Âᵢⱼ = σ(zᵢᵀ zⱼ) - 目标:最小化重构误差(如交叉熵)
✅ 应用于链接预测(预测缺失边)。
6.2 图生成
- 学习整个图的分布
p(G)。 - 方法:变分图自编码器(VGAE)、GAN、扩散模型。
七、实战1:使用PyG进行节点分类(Cora数据集)
7.1 环境准备
pip install torch-scatter torch-sparse torch-cluster torch-spline-conv torch-geometric -f https://data.pyg.org/whl/torch-2.0.0+cpu.html
7.2 加载数据集
import torch
from torch_geometric.datasets import Planetoid
from torch_geometric.nn import GCNConv
import torch.nn.functional as F
# 加载Cora引文网络
dataset = Planetoid(root='/tmp/Cora', name='Cora')
data = dataset[0] # Data(x=[2708, 1433], edge_index=[2, 10556], y=[2708], ...)
print(data)
# 输出:包含特征、边、标签等
7.3 构建GCN模型
class GCN(torch.nn.Module):
def __init__(self, num_features, hidden_dim, num_classes):
super().__init__()
self.conv1 = GCNConv(num_features, hidden_dim)
self.conv2 = GCNConv(hidden_dim, num_classes)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index)
x = F.relu(x)
x = F.dropout(x, training=self.training)
x = self.conv2(x, edge_index)
return F.log_softmax(x, dim=1)
# 模型实例化
model = GCN(dataset.num_features, 16, dataset.num_classes)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
7.4 训练与评估
model.train()
for epoch in range(200):
optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = F.nll_loss(out[data.train_mask], data.y[data.train_mask])
loss.backward()
optimizer.step()
# 评估
model.eval()
pred = model(data.x, data.edge_index).argmax(dim=1)
acc = (pred[data.test_mask] == data.y[data.test_mask]).sum() / data.test_mask.sum()
print(f'测试准确率: {acc:.4f}')
✅ GCN在Cora上可达80%+准确率。
八、实战2:使用GAE进行链接预测
from torch_geometric.nn import GAE, GCNEncoder
from torch_geometric.utils import train_test_split_edges
# 准备数据(划分训练/测试边)
data.train_mask = data.val_mask = data.test_mask = None
data = train_test_split_edges(data)
# 构建GAE模型
class Encoder(torch.nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv1 = GCNConv(in_channels, 2*out_channels)
self.conv2 = GCNConv(2*out_channels, out_channels)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index).relu()
return self.conv2(x, edge_index)
model = GAE(Encoder(dataset.num_features, 16))
# 训练
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
def train():
model.train()
optimizer.zero_grad()
z = model.encode(data.x, data.train_pos_edge_index)
loss = model.recon_loss(z, data.train_pos_edge_index, data.train_neg_adj_mask)
loss.backward()
optimizer.step()
return loss
for epoch in range(100):
loss = train()
# 链接预测
model.eval()
z = model.encode(data.x, data.train_pos_edge_index)
# 计算所有节点对的得分
scores = z @ z.t()
# 预测测试边是否存在
九、GNN的挑战与前沿
9.1 挑战
- ❌ 过平滑(Over-smoothing):层数过深,节点表示趋同。
- ❌ 长距离依赖:信息传递受限于GNN层数。
- ❌ 可扩展性:全图训练内存消耗大。
9.2 前沿方向
- ✅ GraphSAGE:采样邻居,支持归纳学习。
- ✅ GATv2:动态注意力。
- ✅ Graph Transformer:将Transformer应用于图。
- ✅ 3D GNN:用于分子构象预测。
十、总结与学习建议
本文我们:
- 理解了图数据的不规则性;
- 掌握了消息传递范式;
- 学习了GCN与GAT的核心机制;
- 实战了节点分类与链接预测;
- 认识了GAE在图生成中的应用。
📌 学习建议:
- 掌握PyG:最主流的GNN库。
- 理解邻接矩阵:
edge_index是稀疏表示。- 避免过拟合:GNN易过拟合小图。
- 关注采样:GraphSAGE、ClusterGCN解决可扩展性。
- 探索应用:推荐、生物、化学、金融。
十一、下一篇文章预告
第18篇:强化学习基础:Q-Learning与深度Q网络(DQN)
我们将深入讲解:
- 强化学习(RL)的框架(智能体、环境、奖励)
- 马尔可夫决策过程(MDP)
- Q-Learning算法与贝尔曼方程
- 深度Q网络(DQN)与经验回放
- 使用Gym训练智能体玩CartPole游戏
进入“试错学习”的智能世界——强化学习!
参考文献
- Kipf, T. N. & Welling, M. (2016). Semi-Supervised Classification with Graph Convolutional Networks. ICLR.
- Velickovic, P. et al. (2017). Graph Attention Networks. ICLR.
- Hamilton, W. L. (2020). Graph Representation Learning. Synthesis Lectures on Artificial Intelligence and Machine Learning.
- PyTorch Geometric文档: PyG Documentation — pytorch_geometric documentation
更多推荐
所有评论(0)