PyTorch Geometric 数据集快速上手:从下载 Cora 到跑通第一次训练
2026/9/5 16:32:31 网站建设 项目流程

PyTorch Geometric 数据集快速上手:从下载 Cora 到跑通第一次训练

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

PyTorch Geometric(PyG)把上百种图数据集封装成两行 API,但新手第一次导入时,下载卡住、路径报错、缺依赖是三大常客。本文就干一件事:带你把 Cora 数据集从下载到跑通第一次训练,全程不绕路。

两行代码启动:加载、下载、缓存一次讲清

核心 API 就是Planetoid这一个类,Cora、CiteSeer、PubMed 都走它:

from torch_geometric.datasets import Planetoid # root 是本地缓存目录,name 指定具体数据集 dataset = Planetoid(root='data/Planetoid', name='Cora') data = dataset[0] # 取出唯一的图对象 print(dataset.num_features) # 1433,节点特征维度 print(dataset.num_classes) # 7 个类别 print(len(data)) # 2708 个节点

第一次运行:自动下载加处理

  • 构造函数会创建data/Planetoid/raw/,把官方数据包拉下来解压
  • 解析后生成processed/里的 .pt 文件,之后的运行全读这个
  • 第二次运行不再碰网络,直接命中本地缓存

data 对象里有什么

  • data.x:节点特征矩阵,Cora 是 2708×1433 的词袋
  • data.edge_index:边列表,形状 2×E 的长整型张量
  • data.ytrain_mask/val_mask/test_mask:标签和三个互不重叠的划分掩码

跑完这步你会发现,数据集导入其实只发生一次,后面全是本地 IO。真正麻烦在后面——mask 怎么用在训练里,以及网络环境不好时下载怎么救。

接上 GCN:损失降下来的第一次训练

Cora 是单图数据集,dataset[0]拿到的整张图直接喂模型,不需要 DataLoader 分批。先用单层 GCNConv 把链路打通:

import torch.nn.functional as F from torch_geometric.nn import GCNConv model = GCNConv(dataset.num_features, dataset.num_classes) # 单层先跑通 opt = torch.optim.Adam(model.parameters(), lr=0.01) for epoch in range(100): opt.zero_grad() out = model(data.x, data.edge_index) # 一次消息传递 loss = F.cross_entropy(out[data.train_mask], data.y[data.train_mask]) loss.backward(); opt.step() print((out.argmax(-1)[data.val_mask] == data.y[data.val_mask]).float().mean())

mask 是唯一需要记住的东西

  • 训练只用train_mask算损失,Cora 里约 140 个节点
  • 验证、测试各看各的 mask,三者互不重叠
  • 想复现论文精度,换两层 GCN 加 dropout,参考仓库里的examples/gcn.py

特征先归一化

  • 官方示例都带transform=T.NormalizeFeatures(),每列特征除以自己的 L2 范数
  • 不做的后果是损失震荡、收敛慢,词袋特征尤其明显
  • 加在构造函数里即可:Planetoid(root, 'Cora', transform=T.NormalizeFeatures())

训练一百个 epoch,val 准确率在 0.55 上下波动属正常(单层模型的合理水平)。到这里,"下载→导入→训练"的闭环已经完整了。

三个高频卡点逐个拆

下载失败、路径报错、缺依赖,按出现频率排个序。

卡点一:连不上数据源

ConnectionError: HTTPSConnectionPool ... Connection refused

说白了就是网络够不到官方下载地址。

  • 设置代理后重试:os.environ['https_proxy'] = 'http://proxy:port'
  • 手动下载压缩包解压,文件放进data/Planetoid/raw/,重跑会自动跳过下载
  • 公司内网环境优先走第二条,稳定得多

卡点二:目录不存在或没写权限

FileNotFoundError: [Errno 2] No such file or directory: 'data/Planetoid'

说白了就是 root 路径相对的是脚本运行时的工作目录,不是你写代码时所在目录。

  • 用绝对路径:root='/home/user/pyg_data/Planetoid'
  • 先手动mkdir -p建目录,确认当前用户可写
  • 打印os.getcwd()核对脚本真实运行位置

卡点三:稀疏扩展依赖缺失

ModuleNotFoundError: No module named 'torch_sparse'

说白了就是 torch-scatter / torch-sparse 没装,或装的版本和当前 PyTorch、CUDA 对不上。

  • 先升级 PyG 到最新版,新版核心链路已不强依赖这两个包
  • 确实需要时,按官方文档匹配 PyTorch 与 CUDA 版本再编译安装
  • 报错里提到 ABI 的话,说明 wheel 是别的 PyTorch 版本编的,重装即可

不满足于官方数据集:把自定义图喂进 Dataset

自己的数据(比如一张 CSV 边表)不用改 PyG 任何源码,继承Dataset填两个方法就行,骨架比想象中短:

from torch_geometric.data import Dataset class MyDataset(Dataset): def __init__(self, root, transform=None): super().__init__(root, transform) # 建目录、管缓存全在基类 @property def raw_file_names(self): # 文件齐全就自动跳过下载 return ['edges.csv'] def process(self): # 把 raw 解析成 Data 存进 processed/ ...

raw_file_names 是跳过逻辑的开关

  • 基类检查raw/里这些文件是否都在,在就不调download()
  • 手动把数据丢进raw/后重跑,就是"手动下载"的标准姿势

get 方法决定怎么取

  • get(idx)返回第 idx 个图,len()告诉 PyG 一共有几个
  • 单图数据集返回 1 个 Data 即可,多图表数据每个 idx 读一行

大图走磁盘

  • 默认 InMemoryDataset 会把 processed 全读进内存
  • 百万节点级别换OnDiskDataset,按索引动态加载样本
  • 完整定义见torch_geometric/data/dataset.py,骨架和本文代码完全一致

下一步:先用 Planetoid 跑通 Cora 并盯住 val 准确率曲线,再照 MyDataset 骨架把你自己的 CSV 边表喂进去——骨架里真正要你自己写的只有get一个方法。

延伸阅读

  • 自定义数据集教程:docs/source/notes/create_dataset.rst
  • CSV 数据读入示例:docs/source/notes/load_csv.rst
  • 预定义数据集完整列表:torch_geometric/datasets/__init__.py
  • 官方 GCN 训练脚本:examples/gcn.py

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询