摘要:
本文系统讲解图神经网络(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)更新节点表示:

  1. 聚合(Aggregate):收集邻居信息
  2. 更新(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 + I
    • D̃:Ã的度矩阵

✅ GCN是谱图卷积的一阶近似。


五、图注意力网络(GAT)

5.1 核心思想

使用注意力机制为不同邻居分配不同权重。

5.2 计算步骤

  1. 计算注意力系数:

    eᵢⱼ = a(W hᵢ, W hⱼ)
    
    • a:注意力函数(如单层MLP + LeakyReLU)
  2. 归一化(softmax):

    αᵢⱼ = softmaxⱼ(eᵢⱼ) = exp(eᵢⱼ) / Σ_{k∈N(i)} exp(eᵢₖ)
    
  3. 加权求和:

    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在图生成中的应用。

📌 学习建议:

  1. 掌握PyG:最主流的GNN库。
  2. 理解邻接矩阵:edge_index是稀疏表示。
  3. 避免过拟合:GNN易过拟合小图。
  4. 关注采样:GraphSAGE、ClusterGCN解决可扩展性。
  5. 探索应用:推荐、生物、化学、金融。

十一、下一篇文章预告

第18篇:强化学习基础:Q-Learning与深度Q网络(DQN)
我们将深入讲解:

  • 强化学习(RL)的框架(智能体、环境、奖励)
  • 马尔可夫决策过程(MDP)
  • Q-Learning算法与贝尔曼方程
  • 深度Q网络(DQN)与经验回放
  • 使用Gym训练智能体玩CartPole游戏

进入“试错学习”的智能世界——强化学习!


参考文献

  1. Kipf, T. N. & Welling, M. (2016). Semi-Supervised Classification with Graph Convolutional Networks. ICLR.
  2. Velickovic, P. et al. (2017). Graph Attention Networks. ICLR.
  3. Hamilton, W. L. (2020). Graph Representation Learning. Synthesis Lectures on Artificial Intelligence and Machine Learning.
  4. PyTorch Geometric文档: PyG Documentation — pytorch_geometric documentation

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐