Graph Neural Network 图神经网络入门

FreeGuideOnline 最新 2026-07-09

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) 的统一范式。每一层执行三个步骤:

  1. 消息计算:对每条边,生成从源节点发送到目标节点的消息。
  2. 聚合:目标节点收集所有来自邻居的消息。
  3. 更新:结合自身特征与聚合消息,更新节点表示。

数学形式:设 ( 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) 的张量,存储每条边的源节点和目标节点。
  • 模型自动处理消息传递,只需关心层定义。

常见应用场景

场景 任务 示例
节点分类 预测节点属性 社交网络中用户兴趣标签
链接预测 预测缺失边或未来连接 推荐系统中的友邻推荐
图分类 对整张图打标签 分子性质预测、蛋白质功能分类
图生成 生成具有特定属性的图 药物分子设计
知识图谱推理 推断实体间缺失关系 问答系统、搜索引擎

学习路径与进阶资源

  1. 数学基础:线性代数、概率论、图论基本概念。
  2. 经典论文
    • 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).
  3. 工具库:PyTorch Geometric, Deep Graph Library (DGL), Spektral.
  4. 进阶主题
    • 谱方法与空间方法对比
    • 过平滑问题及缓解(如 DropEdge、PairNorm)
    • 大规模图训练(采样、分布式)
    • 异构图、时空图、动态图

总结

图神经网络为结构化数据建模提供了强大且灵活的框架。从 GCN 的简单归一化聚合,到 GAT 的注意力机制,再到 GraphSAGE 的归纳能力,各种变体在不同任务中展现了卓越性能。掌握消息传递的基本范式后,你可以轻松理解大多数 GNN 架构,并应用于社交分析、药物发现、推荐系统等现实问题。

建议边学边实践,在 Cora、CiteSeer、PubMed 等标准数据集上复现模型,并逐步尝试更复杂的图数据,以加深对 GNN 工作原理的理解。