☰
PyTorch Geometric 2.4 实战:3步构建GNN节点分类模型(Cora数据集)
2026/10/10 5:28:03 网站建设 项目流程

PyTorch Geometric 2.4实战:3步构建Cora数据集节点分类模型

为什么选择图神经网络处理Cora数据集?

学术文献引用网络是典型的图结构数据——论文作为节点,引用关系作为边。Cora数据集包含2708篇机器学习论文,每个节点具有1433维的词袋特征向量,标签是论文所属的7个类别之一。传统CNN/RNN难以直接处理这种非欧几里得数据,而PyTorch Geometric(PyG)作为专门处理图数据的深度学习库,提供了高效的图卷积操作实现。

最近在KDD 2023上发布的PyG 2.4版本带来了多项性能优化:

  • 稀疏矩阵运算速度提升40%
  • 新增异构图支持
  • 内存占用降低30%
import torch from torch_geometric.datasets import Planetoid # 自动下载并加载Cora数据集 dataset = Planetoid(root='/tmp/Cora', name='Cora') data = dataset[0] # 获取图数据对象 print(f"节点数量: {data.num_nodes}") print(f"边数量: {data.num_edges}") print(f"特征维度: {data.num_node_features}") print(f"类别数: {dataset.num_classes}")

模型构建与训练实战

1. 数据准备与预处理

Cora数据集已内置在PyG中,但实际项目中常需要自定义图数据。PyG使用Data对象表示图:

from torch_geometric.data import Data # 手动构建图数据的示例 edge_index = torch.tensor([[0, 1, 1, 2], [1, 0, 2, 1]], dtype=torch.long) x = torch.randn(3, 16) # 3个节点,每个16维特征 data = Data(x=x, edge_index=edge_index)

提示:Cora数据集中的边是无向的,但存储时只保存了单向边。实际使用时需要添加反向边:

edge_index = torch.cat([data.edge_index, data.edge_index.flip(0)], dim=1)

2. 构建GNN模型

PyG 2.4提供了多种图卷积层,我们构建一个包含GCNConv的简单网络:

import torch.nn.functional as F from torch_geometric.nn import GCNConv 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 F.log_softmax(x, dim=1)

模型参数对比:

参数名作用设置建议
in_channels输入特征维度与数据集特征维度一致
hidden_channels隐藏层维度通常16-256之间
out_channels输出类别数与数据集类别数一致

3. 训练与评估

PyG的训练循环与常规PyTorch类似,但需注意:

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 train(): model.train() 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() return loss.item() def test(): model.eval() out = model(data.x, data.edge_index) pred = out.argmax(dim=1) accs = [] for mask in [data.train_mask, data.val_mask, data.test_mask]: accs.append(int((pred[mask] == data.y[mask]).sum()) / int(mask.sum())) return accs for epoch in range(1, 201): loss = train() train_acc, val_acc, test_acc = test() if epoch % 50 == 0: print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}, ' f'Train: {train_acc:.4f}, Val: {val_acc:.4f}, ' f'Test: {test_acc:.4f}')

性能优化技巧

1. 消息传递优化

PyG 2.4的消息传递机制进行了重构:

from torch_geometric.nn import MessagePassing class CustomConv(MessagePassing): def __init__(self, in_channels, out_channels): super().__init__(aggr='add') # 支持mean, max, add等聚合方式 self.lin = torch.nn.Linear(in_channels, out_channels) def forward(self, x, edge_index): return self.propagate(edge_index, x=x) def message(self, x_j): return self.lin(x_j)

2. 采样与批处理

对于大规模图,可使用邻居采样:

from torch_geometric.loader import NeighborLoader loader = NeighborLoader( data, num_neighbors=[25, 10], # 两层采样,每层采样数 batch_size=32, input_nodes=data.train_mask ) for batch in loader: train_on_batch(batch)

3. 混合精度训练

PyG 2.4全面支持AMP自动混合精度:

scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): out = model(data.x, data.edge_index) loss = F.nll_loss(out[data.train_mask], data.y[data.train_mask]) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

模型解释与可视化

理解GNN的决策过程至关重要:

import networkx as nx import matplotlib.pyplot as plt from torch_geometric.utils import to_networkx # 将子图转换为NetworkX格式 subgraph = data.subgraph(data.test_mask) G = to_networkx(subgraph, to_undirected=True) # 绘制节点分类结果 plt.figure(figsize=(12, 8)) nx.draw_spring(G, node_color=subgraph.y.cpu().numpy(), cmap='Set2', with_labels=False, node_size=50) plt.show()

节点重要性分析:

from captum.attr import IntegratedGradients ig = IntegratedGradients(model) attr, _ = ig.attribute(data.x, target=data.y, additional_forward_args=(data.edge_index,), return_convergence_delta=True)

进阶应用方向

1. 异构图处理

PyG 2.4增强了对异构图的支持:

from torch_geometric.data import HeteroData data = HeteroData() data['user'].x = torch.randn(100, 32) # 100个用户 data['item'].x = torch.randn(50, 32) # 50个商品 data['user', 'rates', 'item'].edge_index = torch.tensor([[0, 1], [0, 2]])

2. 动态图神经网络

处理时序图数据:

from torch_geometric.nn import TGCN class TemporalGNN(torch.nn.Module): def __init__(self, in_channels): super().__init__() self.tgnn = TGCN(in_channels, 64) self.linear = torch.nn.Linear(64, dataset.num_classes) def forward(self, x, edge_index, edge_weight): h = self.tgnn(x, edge_index, edge_weight) h = F.relu(h) h = self.linear(h) return h

3. 图自监督学习

from torch_geometric.nn import 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) encoder = Encoder(dataset.num_features, 32) model = GAE(encoder) # 图自编码器

常见问题排查

  1. 内存不足:

    • 使用NeighborLoader进行采样
    • 启用torch.use_deterministic_algorithms(False)
  2. 梯度爆炸:

    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
  3. 过拟合:

    • 增加Dropout率
    • 添加L2正则化
    • 使用更小的隐藏层维度
  4. 性能瓶颈分析:

    with torch.profiler.profile( activities=[torch.profiler.ProfilerActivity.CUDA], record_shapes=True) as prof: model(data.x, data.edge_index) print(prof.key_averages().table(sort_by="cuda_time_total"))

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询