PyG 自定义图数据集完全指南:Dataset 与 InMemoryDataset 的实现原理与实战
2026/9/12 16:28:27 网站建设 项目流程

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.Datasettorch_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_dirroot/raw存放下载得到的原始数据
processed_dirroot/processed存放经过process处理后的数据

这两个属性在源码中是直接由root拼接出来的:raw_dir返回osp.join(self.root, 'raw')processed_dir返回osp.join(self.root, 'processed')。构造时若传入的是字符串路径,还会经过osp.expanduser(fs.normpath(root))规范化处理,支持~之类的路径写法。

对应的还有两组文件路径属性:

  • raw_pathsraw_file_names中每个文件名与raw_dir拼接后的绝对路径列表;
  • processed_pathsprocessed_file_names中每个文件名与processed_dir拼接后的绝对路径列表。

它们由raw_file_names/processed_file_names属性驱动,在 dataset.py 中统一实现。因此你只需要实现raw_file_namesprocessed_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_transformpre_filter的字符串表示分别保存为processed_dir/pre_transform.ptprocessed_dir/pre_filter.pt(见 dataset.py)。下次实例化时,如果检测到传入的pre_transform/pre_filter与已保存的不一致,会打印警告,提示你如果确实要更换预处理方式,需要显式传force_reload=True重新处理。

创建内存数据集(InMemoryDataset)

要让一个类成为InMemoryDataset,需要实现四个核心成员(前两个是 property,后两个是方法):

成员类型作用
raw_file_namesproperty返回raw_dir中必须存在的文件列表,用于判断是否可以跳过下载
processed_file_namesproperty返回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=Falseadd_batch=False,不会附加batch向量,而是把区间信息记入slices

最后,在__init__中需要把这两样东西加载为self.dataself.slices两个属性,供get(idx)按索引还原样本。

PyG >= 2.4 的变化:save / load 统一接口

原教程特别提示:从 PyG 2.4 起torch.savecollate的功能被统一封装到InMemoryDataset.save之后,self.dataself.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_urlextract_*process内则是从self.raw_paths读取原始文件、解析出xedge_indexy等张量并构造Data对象。

真实数据集参考:Flickr 与 KarateClub

仓库中的内置数据集是学习自定义实现的最佳范本。

torch_geometric/datasets/flickr.py 展示了最完整的形态:raw_file_names返回 4 个原始文件,downloaddownload_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则完整走processsaveload的落盘流程,二者共同验证了两种构建方式都能正确还原出num_nodesxedge_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 的执行?

只需要不覆写downloadprocess方法即可。PyG 在构造时会通过overrides_method检测你的类是否真正定义了这两个方法(见 dataset.py 与 dataset.py):has_downloadFalse就不执行_download()has_processFalse就不执行_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_dirprocessed_file_names中第一个文件名拼接后的绝对路径)。依据:processed_dir返回osp.join(root, 'processed')processed_pathsprocessed_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则读取该文件并把dataslices分别恢复为self.dataself.slices,此后即可通过dataset[i]按索引还原任意单个样本。

总结

自定义图数据集的核心可归纳为一张"分工表":InMemoryDataset适合整体入内存的数据,实现raw_file_names/processed_file_names/download/process四个成员,靠save(内部collate)与load完成高效存取;Dataset适合超大规模数据,额外实现lenget逐图加载;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),仅供参考

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

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

立即咨询