PyG 自定义图数据集完全指南:Dataset 与 InMemoryDataset 的实现原理与实战
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
本教程讲解如何在 PyTorch Geometric(PyG)中创建属于自己的图数据集。PyG 内置了大量开箱即用的数据集,但在处理自采数据或非公开数据时,你依然需要自己实现数据集类。读完本文,你将掌握torch_geometric.data.Dataset与torch_geometric.data.InMemoryDataset两个抽象类的完整用法、目录约定、transform/pre_transform/pre_filter三种钩子函数的差异,并能从零实现一个可下载、可缓存、可被DataLoader直接消费的自定义图数据集。本文以官方教程 docs/source/tutorial/create_dataset.rst 为主线,结合仓库源码逐层剖析底层机制。
两个抽象类:Dataset 与 InMemoryDataset
PyG 为自定义数据集提供了两个抽象基类,均位于 torch_geometric/data/init.py 的导出列表中:
torch_geometric.data.Dataset:通用数据集基类,继承自torch.utils.data.Dataset,适合无法整体放入内存的大规模图数据;torch_geometric.data.InMemoryDataset:继承自Dataset,在基类之上实现了"把整个数据集一次性加载进 CPU 内存"的缓存机制,适合中小规模、能放进内存的数据集。
从源码看,torch_geometric/data/in_memory_dataset.py 中class InMemoryDataset(Dataset)直接继承基类,两者的__init__签名完全一致(在 dataset.py 中定义):
def __init__( self, root: Optional[str] = None, transform: Optional[Callable] = None, pre_transform: Optional[Callable] = None, pre_filter: Optional[Callable] = None, log: bool = True, force_reload: bool = False, ) -> None:除教程重点介绍的四个参数外,log控制是否在下载/处理时打印控制台输出(默认True),force_reload用于强制重新处理数据集(默认False),后者在后文"跳过 download/process"一节会再次提及。
目录约定:root、raw_dir 与 processed_dir
遵循torchvision的惯例,每个数据集在构造时接收一个root文件夹参数,用于指明数据集的存储位置。PyG 会把root拆分为两个子目录(见 dataset.py):
| 属性 | 路径 | 用途 |
|---|---|---|
raw_dir | root/raw | 存放下载得到的原始数据 |
processed_dir | root/processed | 存放经过process处理后的数据 |
这两个属性在源码中是直接由root拼接出来的:raw_dir返回osp.join(self.root, 'raw'),processed_dir返回osp.join(self.root, 'processed')。构造时若传入的是字符串路径,还会经过osp.expanduser(fs.normpath(root))规范化处理,支持~之类的路径写法。
对应的还有两组文件路径属性:
raw_paths:raw_file_names中每个文件名与raw_dir拼接后的绝对路径列表;processed_paths:processed_file_names中每个文件名与processed_dir拼接后的绝对路径列表。
它们由raw_file_names/processed_file_names属性驱动,在 dataset.py 中统一实现。因此你只需要实现raw_file_names和processed_file_names两个 property,PyG 会自动为你拼好完整路径。
三个钩子函数:transform、pre_transform 与 pre_filter
每个数据集构造时都可以传入三个可选函数,默认均为None。三者的执行时机和用途有本质区别,务必区分清楚:
transform:在每次访问数据对象时动态地对数据做变换(源码中Dataset.__getitem__在取到数据后执行data = self.transform(data),见 dataset.py)。因为每次访问都会执行,所以它最适合数据增强(data augmentation)这类需要"每次不同"的变换,例如随机翻转、随机扰动;pre_transform:在数据保存到磁盘之前执行一次变换,适合只计算一次的重型预处理(如计算图拉普拉斯、构造SparseTensor、特征归一化)。由于处理结果会被固化到processed_dir,第二次实例化数据集时不会再执行;pre_filter:在数据保存之前手动过滤数据对象。它的签名是"输入一个Data对象,返回布尔值",返回False的数据会被丢弃。典型场景是只保留特定类别的样本。
三种函数的完整签名与语义在 in_memory_dataset.py 与 dataset.py 的 docstring 中有精确定义。
这里有一个容易忽略的细节:_process在执行时会额外把pre_transform与pre_filter的字符串表示分别保存为processed_dir/pre_transform.pt和processed_dir/pre_filter.pt(见 dataset.py)。下次实例化时,如果检测到传入的pre_transform/pre_filter与已保存的不一致,会打印警告,提示你如果确实要更换预处理方式,需要显式传force_reload=True重新处理。
创建内存数据集(InMemoryDataset)
要让一个类成为InMemoryDataset,需要实现四个核心成员(前两个是 property,后两个是方法):
| 成员 | 类型 | 作用 |
|---|---|---|
raw_file_names | property | 返回raw_dir中必须存在的文件列表,用于判断是否可以跳过下载 |
processed_file_names | property | 返回processed_dir中必须存在的文件列表,用于判断是否可以跳过处理 |
download | 方法 | 将原始数据下载到raw_dir |
process | 方法 | 读取原始数据、构造Data对象列表,并保存到processed_dir |
下载和解压可以借助torch_geometric.data中现成的工具函数(见 torch_geometric/data/download.py 与 torch_geometric/data/extract.py):
download_url(url, folder, log=True, filename=None):从 URL 下载文件到指定目录,已存在的同名文件会直接复用(打印Using existing file ...并返回路径),下载时会自动创建目录;download_google_url(id, folder, filename, log=True):通过 Google Drive 文件 ID 下载;extract_tar(path, folder, mode='r:gz')、extract_zip(path, folder)、extract_bz2(path, folder)、extract_gz(path, folder):解压各类压缩包到指定目录。
这些函数都通过 torch_geometric/data/init.py 导出,因此可以直接from torch_geometric.data import download_url, extract_zip使用。
process 的核心:collate 与 save / load
process的魔法在于:我们读取原始数据后需要构造一个Data对象列表并保存到processed_dir。如果直接序列化一个巨大的 Python 列表,速度会很慢。因此 PyG 通过collate机制先把列表合并(collate)成一个巨大的Data对象再保存:
- 合并后的大对象把所有样本拼接在一起(各属性的拼接维度由
Data.__cat_dim__决定),同时返回一个slices字典; slices记录了每个样本在每个属性中的起止区间,用于从大对象中还原出任意单个样本。
从 in_memory_dataset.py 的源码可以看到collate的签名:
@staticmethod def collate(data_list): r"""Collates a list of Data or HeteroData objects to the internal storage format of InMemoryDataset.""" if len(data_list) == 1: return data_list[0], None data, slices, _ = collate( data_list[0].__class__, data_list=data_list, increment=False, add_batch=False, ) return data, slices注意collate的底层实现位于 torch_geometric/data/collate.py,它与DataLoader批量打包共享同一套拼接逻辑;区别在于数据集这里increment=False、add_batch=False,不会附加batch向量,而是把区间信息记入slices。
最后,在__init__中需要把这两样东西加载为self.data和self.slices两个属性,供get(idx)按索引还原样本。
PyG >= 2.4 的变化:save / load 统一接口
原教程特别提示:从 PyG 2.4 起,torch.save与collate的功能被统一封装到InMemoryDataset.save之后,self.data和self.slices的加载也被封装到InMemoryDataset.load中。
对照源码(in_memory_dataset.py):
@classmethod def save(cls, data_list, path): """Saves a list of data objects to the file path `path`.""" data, slices = cls.collate(data_list) fs.torch_save((data.to_dict(), slices, data.__class__), path) def load(self, path, data_cls=Data): """Loads the dataset from the file path `path`.""" out = fs.torch_load(path) ... if len(out) == 2: # Backward compatibility. data, self.slices = out else: data, self.slices, data_cls = out if not isinstance(data, dict): # Backward compatibility. self.data = data else: self.data = data_cls.from_dict(data)可见save内部先调用collate得到(data, slices),随后以(data.to_dict(), slices, data.__class__)三元组形式保存;load则兼容两种旧格式(长度为 2 的元组、非 dict 的data),能够平滑读取 PyG 2.4 之前生成的缓存文件。
完整示例:MyOwnDataset
把上述要点串起来,一个标准的内存数据集实现如下(来自原教程,代码可直接运行):
import torch from torch_geometric.data import InMemoryDataset, download_url class MyOwnDataset(InMemoryDataset): def __init__(self, root, transform=None, pre_transform=None, pre_filter=None): super().__init__(root, transform, pre_transform, pre_filter) self.load(self.processed_paths[0]) # For PyG<2.4: # self.data, self.slices = torch.load(self.processed_paths[0]) @property def raw_file_names(self): return ['some_file_1', 'some_file_2', ...] @property def processed_file_names(self): return ['data.pt'] def download(self): # Download to `self.raw_dir`. download_url(url, self.raw_dir) ... def process(self): # Read data into huge `Data` list. data_list = [...] if self.pre_filter is not None: data_list = [data for data in data_list if self.pre_filter(data)] if self.pre_transform is not None: data_list = [self.pre_transform(data) for data in data_list] self.save(data_list, self.processed_paths[0]) # For PyG<2.4: # torch.save(self.collate(data_list), self.processed_paths[0])代码中的省略号...表示你需要根据数据格式补齐的部分:download内通常组合使用download_url/download_google_url与extract_*;process内则是从self.raw_paths读取原始文件、解析出x、edge_index、y等张量并构造Data对象。
真实数据集参考:Flickr 与 KarateClub
仓库中的内置数据集是学习自定义实现的最佳范本。
torch_geometric/datasets/flickr.py 展示了最完整的形态:raw_file_names返回 4 个原始文件,download用download_google_url逐个下载,process中解析npz/npy/json文件构造Data(含train_mask/val_mask/test_mask),应用pre_transform后调用self.save([data], self.processed_paths[0]),构造函数末尾调用self.load(self.processed_paths[0])。结构与上文示例一一对应。
torch_geometric/datasets/karate.py 则展示了另一种常见形态——不落盘、直接在内存中构造:其构造函数传入super().__init__(None, transform)(root=None),然后手动构造Data并调用self.data, self.slices = self.collate([data]),等效于把save+load两步合并在内存中完成。
对应的单元测试在 test/data/test_dataset.py:MyTestDataset使用collate内存构建,MyStoredTestDataset则完整走process→save→load的落盘流程,二者共同验证了两种构建方式都能正确还原出num_nodes、x、edge_index等属性。
使用注意:不要直接修改 self.data
源码为InMemoryDataset.data属性设置了警告机制(in_memory_dataset.py):直接访问内部存储格式data会打印提示,建议改用dataset._data访问内部存储,或通过dataset.{attr_name}直接获取所有图的某个属性堆叠结果。原因是直接修改data不会反映到已经缓存的_data_list中,容易引入隐蔽 bug。日常使用只需通过索引dataset[i]访问单个样本即可。
创建大规模数据集(Dataset)
当数据集无法整体放入内存时,使用基类Dataset。它紧密跟随torchvision数据集的概念,在四个成员之外额外要求实现两个方法:
| 成员 | 作用 |
|---|---|
len() | 返回数据集中的样本数量 |
get(idx) | 实现加载单个图的逻辑 |
内部机制上,Dataset.__getitem__会调用self.get(self.indices()[idx])获取数据对象,并在transform非空时对其应用变换(见 dataset.py)。也就是说,你只需要告诉 PyG "怎么取第 i 个图"和"一共有几个图",其余索引、切片、迭代、transform应用都由基类完成。切片索引(如dataset[2:5]、dataset[:0.9])、长整型/布尔型 Tensor 索引、shuffle()等能力均已在基类中实现(见 dataset.py)。
完整示例:逐图保存的 MyOwnDataset
对于大规模数据集,通常在process中逐图保存,在get中逐图加载:
import os.path as osp import torch from torch_geometric.data import Dataset, download_url class MyOwnDataset(Dataset): def __init__(self, root, transform=None, pre_transform=None, pre_filter=None): super().__init__(root, transform, pre_transform, pre_filter) @property def raw_file_names(self): return ['some_file_1', 'some_file_2', ...] @property def processed_file_names(self): return ['data_1.pt', 'data_2.pt', ...] def download(self): # Download to `self.raw_dir`. path = download_url(url, self.raw_dir) ... def process(self): idx = 0 for raw_path in self.raw_paths: # Read data from `raw_path`. data = Data(...) if self.pre_filter is not None and not self.pre_filter(data): continue if self.pre_transform is not None: data = self.pre_transform(data) torch.save(data, osp.join(self.processed_dir, f'data_{idx}.pt')) idx += 1 def len(self): return len(self.processed_file_names) def get(self, idx): data = torch.load(osp.join(self.processed_dir, f'data_{idx}.pt')) return data这里每个图的数据对象在process中被单独保存为一个.pt文件,并在get中按索引手动加载——这正是"不把整个数据集放进内存"的关键:任意时刻内存中只保留一个图。同时注意len()的返回值要与processed_file_names的数量对应,否则索引会越界。
对于规模在内存可承受范围内的数据,教程与源码都更推荐优先使用InMemoryDataset,因为它通过collate把数据压缩成单个张量化的Data对象,访问速度更快。若你的数据量大到内存放不下,再用Dataset逐图加载;此外InMemoryDataset还提供了to_on_disk_dataset()方法(见 in_memory_dataset.py),可将其转换为基于 SQLite 等后端、逐条落盘的OnDiskDataset,用于分布式训练或共享内存受限的场景。
常见问题(FAQ)
如何跳过 download 和/或 process 的执行?
只需要不覆写download和process方法即可。PyG 在构造时会通过overrides_method检测你的类是否真正定义了这两个方法(见 dataset.py 与 dataset.py):has_download为False就不执行_download(),has_process为False就不执行_process()。同时,即使定义了方法,只要raw_paths/processed_paths中的文件已全部存在(files_exist判定),对应流程也会被自动跳过。
class MyOwnDataset(Dataset): def __init__(self, transform=None, pre_transform=None): super().__init__(None, transform, pre_transform)这种"不覆写即跳过"的约定非常实用:例如你想从内存中的Data列表直接构造数据集、完全不需要磁盘 IO 时,就可以让类只实现processed_file_names(甚至也可以跳过),并把数据在__init__中通过collate直接设置。KarateClub 正是这种做法的官方实例。
我真的必须使用这些数据集接口吗?
不需要。与原生 PyTorch 一样,PyG 并不强制你使用Dataset/InMemoryDataset——例如当你想要在飞行中(on the fly)生成合成数据、又不想显式保存到磁盘时,直接构造一个由torch_geometric.data.Data对象组成的普通 Python 列表,丢给torch_geometric.loader.DataLoader即可:
from torch_geometric.data import Data from torch_geometric.loader import DataLoader data_list = [Data(...), ..., Data(...)] loader = DataLoader(data_list, batch_size=32)DataLoader会自动把列表中的多个Data对象拼接(collate)成Batch。需要注意的是,这种方法失去了磁盘缓存、pre_transform只算一次、pre_filter过滤等特性,因此更适合数据量小或数据完全由程序生成的场景。
小练习与解答
原教程给出了一段从Data列表构造InMemoryDataset的示例,请先自行思考再对照解答:
class MyDataset(InMemoryDataset): def __init__(self, root, data_list, transform=None): self.data_list = data_list super().__init__(root, transform) self.load(self.processed_paths[0]) @property def processed_file_names(self): return 'data.pt' def process(self): self.save(self.data_list, self.processed_paths[0])1.self.processed_paths[0]的输出是什么?
是root/processed/data.pt(即processed_dir与processed_file_names中第一个文件名拼接后的绝对路径)。依据:processed_dir返回osp.join(root, 'processed'),processed_paths把processed_file_names的每个文件名与processed_dir拼接(见 dataset.py 与 dataset.py)。这里的processed_file_names返回的是单个字符串'data.pt'而非列表,PyG 的to_list工具会自动把它包装成单元素列表。
2.InMemoryDataset.save做了什么?
它先把data_list通过collate合并成单个Data对象并生成slices字典,然后以(data.to_dict(), slices, data.__class__)的格式序列化到指定路径(见 in_memory_dataset.py)。load则读取该文件并把data与slices分别恢复为self.data与self.slices,此后即可通过dataset[i]按索引还原任意单个样本。
总结
自定义图数据集的核心可归纳为一张"分工表":InMemoryDataset适合整体入内存的数据,实现raw_file_names/processed_file_names/download/process四个成员,靠save(内部collate)与load完成高效存取;Dataset适合超大规模数据,额外实现len与get逐图加载;transform用于每次访问时动态变换(数据增强),pre_transform用于保存前的一次性重型预处理,pre_filter用于保存前的样本过滤;不需要持久化时,直接用Data列表 +DataLoader即可。想深入研究实现细节,建议通读 torch_geometric/data/dataset.py 与 torch_geometric/data/in_memory_dataset.py,并参照 torch_geometric/datasets/flickr.py、torch_geometric/datasets/karate.py 以及 test/data/test_dataset.py 中的测试用例,它们是这两个抽象类最权威的用法示范。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考