Graph Neural Network 图神经网络入门
Graph Neural Network 图神经网络入门
什么是图神经网络
图神经网络(Graph Neural Network,GNN)是一类专门用于处理图结构数据的深度学习模型。与图像(规则网格)或文本(序列)不同,图中的实体通过任意连接关系组织,具有非欧几里得特性,传统的卷积或循环网络无法直接应用。
GNN 的核心思想是:利用节点间的连接关系,将信息沿着图的边进行传播,让每个节点能够聚合邻居节点的特征,从而学习到包含结构信息的表示。这种表示可以用于节点分类、链接预测、图分类等任务。
图的基础概念
在深入 GNN 之前,需要统一图的术语。
| 术语 | 定义 |
|---|---|
| 图 ( G = (V, E) ) | 顶点集合 ( V ),边集合 ( E ) |
| 邻接矩阵 ( A ) | ( A_{ij} = 1 ) 若存在边 ( i \to j ),否则 0 |
| 节点特征 ( X ) | 每个节点 ( v ) 的特征向量 ( x_v \in \mathbb{R}^d ) |
| 度矩阵 ( D ) | 对角阵,( D_{ii} = \sum_j A_{ij} ) |
| 邻居集合 ( \mathcal{N}(v) ) | 与节点 ( v ) 直接相连的节点集 |
图类型举例
- 社交网络:用户为节点,关注关系为边。
- 分子结构:原子为节点,化学键为边。
- 知识图谱:实体为节点,关系为边。
- 推荐系统:用户与物品为节点,交互为边。
为什么需要图神经网络
传统深度学习的局限性:
- CNN 假设输入具有平移不变性的网格结构。
- RNN/Transformer 假设输入是序列。
- 图数据中节点间的关系不规整,邻居数量可变,且顺序无关。
GNN 必须具备以下性质:
- 排列不变性:对节点重新编号后,计算结果不变。
- 局部结构感知:能捕捉以节点为中心的邻域模式。
消息传递框架
现代 GNN 大多遵循消息传递神经网络(MPNN) 的统一范式。每一层执行三个步骤:
- 消息计算:对每条边,生成从源节点发送到目标节点的消息。
- 聚合:目标节点收集所有来自邻居的消息。
- 更新:结合自身特征与聚合消息,更新节点表示。
数学形式:设 ( h_v^{(k)} ) 为节点 ( v ) 在第 ( k ) 层的隐状态,
[ m_{vu}^{(k)} = \text{MESSAGE}\left( h_v^{(k-1)}, h_u^{(k-1)}, e_{vu} \right) ] [ a_v^{(k)} = \text{AGGREGATE}\left( { m_{uv}^{(k)} : u \in \mathcal{N}(v) } \right) ] [ h_v^{(k)} = \text{UPDATE}\left( h_v^{(k-1)}, a_v^{(k)} \right) ]
其中 ( e_{vu} ) 是可选边特征。初始 ( h_v^{(0)} = x_v )。
关键要点
- 聚合函数需满足排列不变性,常见选项:求和、均值、最大值。
- 消息可包含边特征,使模型适应更丰富的关系。
经典 GNN 模型
1. 图卷积网络(GCN)
GCN 定义了最简单有效的图卷积操作。对于第 ( k+1 ) 层,节点更新公式为:
[ H^{(k+1)} = \sigma\left( \tilde{D}^{-\frac{1}{2}} \tilde{A} \tilde{D}^{-\frac{1}{2}} H^{(k)} W^{(k)} \right) ]
其中:
- ( \tilde{A} = A + I )(添加自环的邻接矩阵)
- ( \tilde{D}{ii} = \sum_j \tilde{A}{ij} )
- ( W^{(k)} ) 是可学习权重矩阵
- ( \sigma ) 是非线性激活(如 ReLU)
直观理解:对邻居表示进行对称归一化求和,然后线性变换。
特点:参数少,性能扎实,但不区分邻居重要性,且面对异配图时效果有限。
2. 图注意力网络(GAT)
GAT 引入注意力机制,为不同邻居分配不同的重要性权重。计算流程:
- 注意力系数: [ \alpha_{vu} = \frac{\exp\left(\text{LeakyReLU}\left( \mathbf{a}^\top [W h_v | W h_u] \right)\right)}{\sum_{k \in \mathcal{N}(v) \cup {v}} \exp\left(\text{LeakyReLU}\left( \mathbf{a}^\top [W h_v | W h_k] \right)\right)} ]
- 节点更新: [ h_v' = \sigma\left( \sum_{u \in \mathcal{N}(v) \cup {v}} \alpha_{vu} W h_u \right) ]
其中 ( \mathbf{a} ) 是注意力向量,( | ) 表示拼接。
多头注意力:并行计算多组注意力,再将结果拼接或平均,增强稳定性。
优势:动态权重,适用于异构图,可解释性更高。
3. GraphSAGE
GraphSAGE 设计了多种聚合器(均值、池化、LSTM),以支持归纳式学习(在未见过节点上生成嵌入)。核心方程为:
[ h_v^{(k)} = \sigma\left( W^{(k)} \cdot \text{CONCAT}\left( h_v^{(k-1)}, \text{AGG}\left( { h_u^{(k-1)}, \forall u \in \mathcal{N}(v) } \right) \right) \right) ]
常用聚合器:
- Mean Aggregator:对邻居取均值。
- LSTM Aggregator:对邻居随机排列后输入 LSTM。
- Pooling Aggregator:先对邻居做线性变换+激活,再最大或均值池化。
GraphSAGE 很适合大规模图训练,通过采样邻居子图控制计算代价。
4. 其他重要变体
- GIN(图同构网络):理论证明其判别能力逼近 WL 测试,用于图分类。
- RGCN(关系图卷积网络):处理多种边类型(知识图谱)。
- ChebNet:基于切比雪夫多项式的谱域图卷积。
动手实现:使用 PyTorch Geometric 构建 GCN
PyTorch Geometric (PyG) 是处理图数据的常用库。以下示例展示一个两层 GCN 用于半监督节点分类。
import torch
import torch.nn.functional as F
from torch_geometric.nn import GCNConv
from torch_geometric.datasets import Planetoid
# 加载 Cora 数据集
dataset = Planetoid(root='/tmp/Cora', name='Cora')
data = dataset[0] # 包含 x, edge_index, y, train_mask 等
class GCN(torch.nn.Module):
def __init__(self, in_channels, hidden_channels, out_channels):
super().__init__()
self.conv1 = GCNConv(in_channels, hidden_channels)
self.conv2 = GCNConv(hidden_channels, out_channels)
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 x
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)
model.train()
for epoch in range(200):
optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = F.cross_entropy(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)
correct = (pred[data.test_mask] == data.y[data.test_mask]).sum()
acc = int(correct) / int(data.test_mask.sum())
print(f'Test Accuracy: {acc:.4f}')
关键点:
edge_index是形状(2, num_edges)的张量,存储每条边的源节点和目标节点。- 模型自动处理消息传递,只需关心层定义。
常见应用场景
| 场景 | 任务 | 示例 |
|---|---|---|
| 节点分类 | 预测节点属性 | 社交网络中用户兴趣标签 |
| 链接预测 | 预测缺失边或未来连接 | 推荐系统中的友邻推荐 |
| 图分类 | 对整张图打标签 | 分子性质预测、蛋白质功能分类 |
| 图生成 | 生成具有特定属性的图 | 药物分子设计 |
| 知识图谱推理 | 推断实体间缺失关系 | 问答系统、搜索引擎 |
学习路径与进阶资源
- 数学基础:线性代数、概率论、图论基本概念。
- 经典论文:
- Kipf & Welling, Semi-Supervised Classification with Graph Convolutional Networks (GCN).
- Veličković et al., Graph Attention Networks (GAT).
- Hamilton et al., Inductive Representation Learning on Large Graphs (GraphSAGE).
- 工具库:PyTorch Geometric, Deep Graph Library (DGL), Spektral.
- 进阶主题:
- 谱方法与空间方法对比
- 过平滑问题及缓解(如 DropEdge、PairNorm)
- 大规模图训练(采样、分布式)
- 异构图、时空图、动态图
总结
图神经网络为结构化数据建模提供了强大且灵活的框架。从 GCN 的简单归一化聚合,到 GAT 的注意力机制,再到 GraphSAGE 的归纳能力,各种变体在不同任务中展现了卓越性能。掌握消息传递的基本范式后,你可以轻松理解大多数 GNN 架构,并应用于社交分析、药物发现、推荐系统等现实问题。
建议边学边实践,在 Cora、CiteSeer、PubMed 等标准数据集上复现模型,并逐步尝试更复杂的图数据,以加深对 GNN 工作原理的理解。