PyTorch Geometric 图神经网络实战入门:最短安装路径与最小可运行示例
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
PyTorch Geometric(简称 PyG)是构建在 PyTorch 之上的图神经网络库:它把节点、边、图这些"不规则数据"变成普通张量,让你用十几行代码就能训练 GCN、GAT 等 GNN 模型。读完本文,你会完成从 pip 安装到跑通一个真实数据集(Cora 论文引用图分类)的完整流程,并知道遇到问题时该去仓库哪个目录找答案。
为什么你需要 PyG
如果你做过图像分类,应该知道一张图片就是一个规整的二维张量。但现实世界里大量数据没有这种规整形状:社交网络里谁关注谁、分子中原子如何成键、知识图谱中实体如何关联——它们都是"图"。用普通神经网络直接吃图数据,你得自己处理邻居聚合、变长结构这些麻烦事。PyG 把这些脏活全包了:内置 100 多个数据集、几十种现成的 GNN 层,你只写"把哪一层接哪一层",剩下交给库。
安装与首次验证
最短安装路径就一条命令(前提是已装好 PyTorch):
pip install torch_geometric从 PyG 2.3 起,核心功能不再依赖额外扩展库,装完即可用。装完立刻验证一下,跑这段代码:
import torch from torch_geometric.data import Data x = torch.tensor([[1.0], [2.0], [3.0]]) # 3 个节点的特征 edge_index = torch.tensor([[0, 1], [1, 2]]) # 0->1->2 一条链 data = Data(x=x, edge_index=edge_index) print(data)预期效果:终端会打印出Data(x=[3, 1], edge_index=[2, 2])这样的摘要。看到这个输出,说明"图数据"这个核心对象已经在你的环境里跑起来了。
最小完整示例:三步训练一个 GCN
下面这个例子能完整跑起来:加载数据、定义模型、训练并评估。它做的是"论文分类"——Cora 数据集里每篇论文是一个节点,两篇论文互相引用就有一条边,目标是根据引用关系和论文内容,判断每篇论文属于哪个领域(机器学习、计算机视觉等 7 类)。
生活化类比:想象你要判断一个新朋友属于哪个圈层,最靠谱的方式不是听他自我介绍,而是看他跟谁走得近。GCN 干的就是这件事——"看邻居"。
第一步,加载数据并定义模型:
from torch_geometric.datasets import Planetoid from torch_geometric.nn import GCNConv dataset = Planetoid(root='data/Planetoid', name='Cora') data = dataset[0] class GCN(torch.nn.Module): def __init__(self): super().__init__() self.conv1 = GCNConv(dataset.num_features, 16) self.conv2 = GCNConv(16, dataset.num_classes) def forward(self, x, edge_index): x = self.conv1(x, edge_index).relu() return self.conv2(x, edge_index)第二步,写训练循环——你会发现它和训练任何 PyTorch 模型毫无区别:
import torch.nn.functional as F model = GCN() optimizer = torch.optim.Adam(model.parameters(), lr=0.01) for epoch in range(200): pred = model(data.x, data.edge_index) loss = F.cross_entropy(pred[data.train_mask], data.y[data.train_mask]) optimizer.zero_grad() loss.backward() optimizer.step()训练时只需要train_mask标出的少量节点作为监督信号,验证和测试节点完全留白——这正是半监督节点分类的形态。
第三步,看产出:
model.eval() with torch.no_grad(): pred = model(data.x, data.edge_index).argmax(dim=1) acc = int((pred[data.test_mask] == data.y[data.test_mask]).sum()) acc /= int(data.test_mask.sum()) print(f'Test Accuracy: {acc:.4f}')跑完 200 轮后,测试集准确率通常在 0.80 上下波动,这是 GCN 在 Cora 上的标准水平。想要带调参和评估细节的完整版本,可以直接读仓库里的 examples/gcn.py,它是官方 Quick Tour 同款写法的加强版。
核心概念拆解
看懂上面代码后,只需要再记住下面这些抽象,PyG 的 API 就能自己读了:
| 概念 | 是什么 | 为什么重要 |
|---|---|---|
Data | 单个图数据的容器,核心字段是x(节点特征,形状[num_nodes, num_features])和edge_index(边列表,形状[2, num_edges]) | 几乎所有 API 的输入输出都是它 |
Dataset | 数据集基类,dataset[0]返回一个Data | 内置数据集全部遵循此接口 |
*Conv层 | 如GCNConv、GATConv,输入(x, edge_index),输出新的x | 所有 GNN 模型的积木块 |
transform | 数据预处理函数,如归一化特征 | 在 Dataset 上挂一个钩子,加载时自动应用 |
loader | 如NeighborLoader,对大图按邻居采样出小批量 | 百万节点图装不进显存时用它 |
两点补充说明:
edge_index用的是 COO 坐标格式:第一行是边的起点编号,第二行是终点编号,第i列就描述第i条边。它不保存任何邻接矩阵,所以天然稀疏、省内存。- 消息传递是 PyG 的统一抽象。每一层 Conv 做的事情都是"每个节点收集邻居信息、聚合、更新自己的特征"。读懂 torch_geometric/nn/conv/ 目录下任意一个
*Conv.py,就能理解对应论文的模型结构。
真实场景走通:Cora 上的完整闭环
上一节的代码其实已经把完整闭环走了一遍,这里补充"从输入到输出"的数据流视角,帮你建立全局印象:
- 输入:
Planetoid(root='data/Planetoid', name='Cora')首次运行会自动下载数据到该目录,之后直接读缓存。Cora 有 2708 篇论文、5429 条引用边、1433 维词袋特征、7 个类别。 - 处理:模型两层
GCNConv逐节点聚合邻居特征;数据集自带的train_mask/val_mask/test_mask三个布尔向量划分出 140 / 500 / 1000 个节点,分别用于训练、验证、测试。 - 输出:每个节点一个 7 维 logits 向量,
argmax即预测类别。最终产出就是那个Test Accuracy: 0.8xxx的数字——它意味着模型仅用 140 个带标签节点,就看懂了整张引用网络。
如果换成你的业务,"论文领域分类"可以对应"用户流失预测":节点是用户,边是社交关系,mask 是已知流失/留存的用户。数据结构不变,只是x换成了用户特征。
避坑与进阶
两个新手最高频的坑:
- 训练/评估模式不切换。
forward里用了F.dropout却没在推理前调model.eval(),会导致每次评估结果都不一样。记住"训练循环里model.train(),评估前model.eval()"这个固定搭配,上面示例已经这么写了。 - 边方向写反。
edge_index[0]是起点、edge_index[1]是终点,[[0, 1], [1, 0]]才是无向的 0-1 边,而[[0, 1], [0, 1]]是两条同向边。节点分类任务对方向通常不敏感,但做链接预测时方向就是答案本身,写反会直接归零。
进阶方向按难度排列:
- 大图采样:图大到显存装不下时,用
NeighborLoader按层采样邻居,思路见 examples/reddit.py。 - 异构图:节点和边分多种类型(比如"作者-论文-期刊"),示例集中在 examples/hetero/。
- 点云与 3D 形状:PyG 不只处理 2D 图,PointNet、DGCNN 等点云模型都有现成实现,流程大致是"采样分组 → 局部网络 → 上采样":
- 可解释性:想知道模型为什么给出某个预测,examples/explain/ 里的 GNNExplainer 系列脚本可以直接照抄。
资源入口
- torch_geometric/nn/:所有模型层与模型的源码目录,读代码学模型结构最快
- examples/:按任务分类的示例脚本,基本"改个数据集就能跑"
- torch_geometric/datasets/:内置数据集清单,每个文件就是一个数据集的加载器
- test/:全量单元测试,想看某个 API 的标准用法时,对应的 test 文件就是最好的说明书
装好、跑通 Cora、看懂Data和edge_index,你就已经站在了 PyG 的入门线上。剩下的事情只有一个:挑一个你自己的图数据,把它塞进Data对象里。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考