在 PyTorch 中更轻松地使用图神经网络
大多数开发者习惯于处理清晰、规则的数据结构。图像是像素网格,文本可以轻松转换为 token 序列。常规卷积网络或 Transformer 可以完美处理这些数据。
当数据本质上是非线性时,问题就出现了。比如银行交易、社交网络中的用户关系,或复杂分子的化学式。图结构无处不在。如果尝试将它们输入常规 PyTorch,就必须手动处理稀疏矩阵、编写复杂的节点聚合循环,并追踪索引。
这正是 PyTorch Geometric(简称 PyG)创建的目的。它是 PyTorch 的扩展,为你处理所有底层数学运算,并提供直观的图神经网络 API。
PyG 的结构
如果你已经了解如何在 PyTorch 中编写模型,那么只需几个小时就能熟悉 PyG。该框架遵循相同的原则:相同的模块 torch.nn、熟悉的训练循环,以及直接操作张量的方式。
图网络的所有魔力都围绕 消息传递 概念构建。图中每个顶点从邻居节点收集信息、合并信息,并更新自身状态。
PyG 提供了一个基类 MessagePassing 来抽象这一逻辑。你只需要定义三件事:
- 如何从邻居节点生成消息
- 如何聚合这些消息(求和、平均、最大值)
- 节点本身如何更新
以下是一个简单图卷积层的实现示例:
import torch
from torch.nn import Sequential, Linear, ReLU
from torch_geometric.nn import MessagePassing
class EdgeConv(MessagePassing):
def __init__(self, in_channels, out_channels):
super().__init__(aggr="max")
self.mlp = Sequential(
Linear(2 * in_channels, out_channels),
ReLU(),
Linear(out_channels, out_channels),
)
def forward(self, x, edge_index):
# x задает фичи узлов, edge_index отвечает за связи между ними
return self.propagate(edge_index, x=x)
def message(self, x_j, x_i):
# x_i — текущий узел, x_j — его сосед
edge_features = torch.cat([x_i, x_j - x_i], dim=-1)
return self.mlp(edge_features)
无需编写遍历所有图边的循环。PyG 本身会并行化此操作,并借助编译的 CUDA 内核快速执行。
基础示例:节点分类
让我们使用标准数据集 Cora,它由科学论文及其相互引用关系组成。任务是根据论文文本及其与其他出版物的关联来预测论文类别。
一个常规卷积模型 GCN 配合数据加载,只需二十多行代码即可完成:
import torch
from torch_geometric.datasets import Planetoid
from torch_geometric.nn import GCNConv
# Загружаем датасет
dataset = Planetoid(root='.', name='Cora')
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).relu()
x = self.conv2(x, edge_index)
return x
model = GCN(dataset.num_features, 16, dataset.num_classes)
训练循环对 PyTorch 爱好者来说完全标准:传入特征矩阵 x 和边张量 edge_index,计算 CrossEntropy,然后调用 loss.backward()。
图的主要问题及 PyG 如何解决
图深度学习的主要挑战是扩展性。图像批次可以轻松分割并分批发送到 GPU。但在图中,所有节点都是相互连接的。当图增长到数百万个节点和数十亿条边时,即使是最强大的 GPU 也无法将其全部加载到内存中。
PyG 开发者在解决这个问题上投入了大量努力。该库包含子图采样机制:
NeighborLoader为批次中的每个节点仅选择随机邻居,防止内存爆炸ClusterGCN将大图聚类为独立片段,并在其上训练网络GraphSAINT在保持拓扑结构的同时采样随机子图
得益于此,PyG 可以在普通 GPU 上处理巨型图而不会耗尽内存。
开箱即用的其他功能
该仓库包含大量已实现的研究论文架构。你会发现经典的 GCN、GAT 和 GraphSAGE,以及用于分子分析的专用模型(如 SchNet 和 DimeNet)或用于处理 3D 点云的 PointNet。
除了算法,该库还包含数百个标准数据集的加载器:从社交网络到生物信息学。这节省了大量编写解析器的时间。
安装与陷阱
自 2.3 版本以来,基础包的安装变得非常简单:
pip install torch_geometric
这足以开始使用。但是,如果需要在大型图上获得最大速度,则需要安装带有 C++/CUDA 扩展的附加库:pyg-lib、torch-scatter 和 torch-sparse。
这里有时会变得棘手。这些二进制文件与特定版本的 PyTorch 和 CUDA 紧密耦合。如果安装了不兼容的版本,Python 会抛出大量 C++ 库导入错误。务必查看项目网站上的兼容性矩阵,并使用明确的 wheel URL 安装扩展。
谁适合使用这个框架
如果你的数据不太适合表格或网格,这个库值得一试:
- 推荐系统:
相关项目