PyG 自定义图数据集完全指南:Dataset 与 InMemoryDataset 的实现原理与实战

发布时间:2026/9/12 16:28:29
PyG 自定义图数据集完全指南:Dataset 与 InMemoryDataset 的实现原理与实战 PyG 自定义图数据集完全指南Dataset 与 InMemoryDataset 的实现原理与实战【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric本教程讲解如何在 PyTorch GeometricPyG中创建属于自己的图数据集。PyG 内置了大量开箱即用的数据集但在处理自采数据或非公开数据时你依然需要自己实现数据集类。读完本文你将掌握torch_geometric.data.Dataset与torch_geometric.data.InMemoryDataset两个抽象类的完整用法、目录约定、transform/pre_transform/pre_filter三种钩子函数的差异并能从零实现一个可下载、可缓存、可被DataLoader直接消费的自定义图数据集。本文以官方教程 docs/source/tutorial/create_dataset.rst 为主线结合仓库源码逐层剖析底层机制。两个抽象类Dataset 与 InMemoryDatasetPyG 为自定义数据集提供了两个抽象基类均位于 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控制是否在下载/处理时打印控制台输出默认Trueforce_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_names和processed_file_names两个 propertyPyG 会自动为你拼好完整路径。三个钩子函数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_reloadTrue重新处理。创建内存数据集InMemoryDataset要让一个类成为InMemoryDataset需要实现四个核心成员前两个是 property后两个是方法成员类型作用raw_file_namesproperty返回raw_dir中必须存在的文件列表用于判断是否可以跳过下载processed_file_namesproperty返回processed_dir中必须存在的文件列表用于判断是否可以跳过处理download方法将原始数据下载到raw_dirprocess方法读取原始数据、构造Data对象列表并保存到processed_dir下载和解压可以借助torch_geometric.data中现成的工具函数见 torch_geometric/data/download.py 与 torch_geometric/data/extract.pydownload_url(url, folder, logTrue, filenameNone)从 URL 下载文件到指定目录已存在的同名文件会直接复用打印Using existing file ...并返回路径下载时会自动创建目录download_google_url(id, folder, filename, logTrue)通过 Google Drive 文件 ID 下载extract_tar(path, folder, moder: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 / loadprocess的魔法在于我们读取原始数据后需要构造一个Data对象列表并保存到processed_dir。如果直接序列化一个巨大的 Python 列表速度会很慢。因此 PyG 通过collate机制先把列表合并collate成一个巨大的Data对象再保存合并后的大对象把所有样本拼接在一起各属性的拼接维度由Data.__cat_dim__决定同时返回一个slices字典slices记录了每个样本在每个属性中的起止区间用于从大对象中还原出任意单个样本。从 in_memory_dataset.py 的源码可以看到collate的签名staticmethod def collate(data_list): rCollates 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_listdata_list, incrementFalse, add_batchFalse, ) return data, slices注意collate的底层实现位于 torch_geometric/data/collate.py它与DataLoader批量打包共享同一套拼接逻辑区别在于数据集这里incrementFalse、add_batchFalse不会附加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.pyclassmethod 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_clsData): 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, transformNone, pre_transformNone, pre_filterNone): super().__init__(root, transform, pre_transform, pre_filter) self.load(self.processed_paths[0]) # For PyG2.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 PyG2.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)rootNone然后手动构造Data并调用self.data, self.slices self.collate([data])等效于把saveload两步合并在内存中完成。对应的单元测试在 test/data/test_dataset.pyMyTestDataset使用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, transformNone, pre_transformNone, pre_filterNone): 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, fdata_{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, fdata_{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.pyhas_download为False就不执行_download()has_process为False就不执行_process()。同时即使定义了方法只要raw_paths/processed_paths中的文件已全部存在files_exist判定对应流程也会被自动跳过。class MyOwnDataset(Dataset): def __init__(self, transformNone, pre_transformNone): 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_size32)DataLoader会自动把列表中的多个Data对象拼接collate成Batch。需要注意的是这种方法失去了磁盘缓存、pre_transform只算一次、pre_filter过滤等特性因此更适合数据量小或数据完全由程序生成的场景。小练习与解答原教程给出了一段从Data列表构造InMemoryDataset的示例请先自行思考再对照解答class MyDataset(InMemoryDataset): def __init__(self, root, data_list, transformNone): 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),仅供参考