ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch Geometric 自定义图数据集实战指南:InMemoryDataset 与 Dataset 的创建、加载与扩展

PyTorch Geometric 自定义图数据集实战指南:InMemoryDataset 与 Dataset 的创建、加载与扩展 PyTorch Geometric 自定义图数据集实战指南InMemoryDataset 与 Dataset 的创建、加载与扩展【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric虽然 PyTorch GeometricPyG已经内置了大量高质量数据集如 Planetoid、QM9、OGB 系列等见 torch_geometric/datasets但当你面对自采数据或非公开数据时往往需要亲手实现自己的数据集类。本指南以官方文档 docs/source/tutorial/create_dataset.rst 为核心结合仓库源码系统讲解 PyG 数据集的两大抽象基类InMemoryDataset与Dataset的设计思想、目录约定、四个核心钩子方法与完整实现范例并深入剖析collate/save/load的底层机制。读完本文你将能够独立编写可复现、可缓存、可增量加载的自定义图数据集并与DataLoader无缝衔接。数据集抽象基类两类选择一种约定PyG 在torch_geometric.data命名空间中提供了两个抽象基类参见 torch_geometric/data/init.py 的导出torch_geometric.data.Dataset通用数据集基类继承自torch.utils.data.Dataset适用于无法整体装入内存的大规模数据集torch_geometric.data.InMemoryDataset继承自Dataset当整个数据集能够装入 CPU 内存时应优先选用。两类数据集共享一套统一的目录约定沿用torchvision的惯例每个数据集接收一个root文件夹作为存储根目录并在其下拆分为两个子目录raw_dir默认root/raw存放下载得到的原始数据processed_dir默认root/processed存放加工后的数据。这两个路径分别由 dataset.py 中的raw_dir与processed_dir属性计算得出你也可以像 Planetoid 那样覆写它们以支持按数据集名如root/Cora/raw分目录存放。此外每个数据集都可以接收三个默认均为None的回调函数参数作用时机典型用途transform每次访问数据对象之前动态执行数据增强如随机扰动、掩码pre_transform数据对象保存到磁盘之前执行只需执行一次的重量级预计算如添加虚拟节点、图归一化pre_filter保存前手动过滤数据对象限制数据对象属于特定类别等筛选逻辑三者具体的语义在 in_memory_dataset.py 与 dataset.py 的 docstring 中有完整定义。其中pre_transform与pre_filter的执行结果还会被序列化为processed_dir下的pre_transform.pt与pre_filter.pt文件当再次实例化数据集时如果传入的pre_transform/pre_filter与磁盘上记录的不一致PyG 会发出警告提示你显式传入force_reloadTrue以重新处理见 dataset.py 的_process实现。这一机制有效避免了加工过但钩子函数已变化的静默错误。创建 In-Memory 数据集四个必须实现的方法要创建一个InMemoryDataset你需要实现四个基础方法其抽象定义见 in_memory_dataset.pyraw_file_names属性raw_dir中必须存在的文件列表用于判断是否可以跳过下载processed_file_names属性processed_dir中必须存在的文件列表用于判断是否可以跳过处理download()把原始数据下载到raw_dirprocess()读取原始数据、加工并保存到processed_dir。生命周期下载与处理如何被自动触发Dataset.__init__见 dataset.py在构造时会依次执行if self.has_download: self._download() if self.has_process: self._process()其中has_download/has_process通过overrides_methoddataset.py检测子类是否真正覆写了对应方法_download仅在raw_paths中的文件尚不存在时才调用download()_process则在force_reloadFalse且processed_paths文件已存在时直接跳过处理。这意味着只要raw_file_names/processed_file_names返回的文件都已就位重复实例化数据集不会重复下载或重复处理——这是 PyG 数据集天然具备的缓存能力。完整实现示例官方教程给出了一个最小化但完整的InMemoryDataset实现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])PyG ≥ 2.4 的新机制save / load 取代手动 collate torch.load教程中特别注明从 PyG 2.4 起torch.save与InMemoryDataset.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保存的是三元组(data.to_dict(), slices, data.__class__)load在读取时会优先兼容旧的二元组格式即 PyG 2.4 的(data, slices)并支持自定义data_cls例如HeteroData的反序列化。因此在自定义数据集里__init__结尾调用self.load(self.processed_paths[0])即可完成全部加载。collate 与 slices把一个 Python 列表压缩成一个对象教程强调直接保存一个巨大的 Python 列表非常缓慢因此我们在保存前通过InMemoryDataset.collatein_memory_dataset.py把列表拼接成一个巨大的Data对象并额外得到用于还原单个样本的slices字典。其底层由 torch_geometric/data/collate.py 的collate函数完成将所有样本按属性如x、edge_index、y纵向拼接为统一表示维护slice_dict记录每个属性在各样本间的切片边界用于从大对象中重构单个样本维护inc_dict记录各属性需要累加的量——例如edge_index需要按前序样本的节点数累加偏移见 collate.py 中关于inc_dict的注释还原时再做递减。InMemoryDataset.get(idx)in_memory_dataset.py正是利用separate与slices从self._data中切出第idx个样本并带有一层缓存_data_list与拷贝保护当数据集只有单个样本slices is None时len()返回 1get(0)直接返回_data的浅拷贝。这就是self.data与self.slices两枚属性的全部用途——同时要注意InMemoryDataset.data属性会发出不建议直接访问内部存储的警告in_memory_dataset.py推荐通过dataset[0]或dataset.x等接口访问。仓库中的现成范例KarateClub最简单的InMemoryDataset直接在内存中构造Data(x, edge_index, y, train_mask)最后一行self.data, self.slices self.collate([data])展示了手动 collate 的经典写法Planetoid展示了raw_dir/processed_dir覆写、raw_file_names返回带前缀的文件名列表ind.cora.x等、download中调用download_url拉取远程数据以及在__init__里根据split参数public/full/random等对self.data, self.slices self.collate([data])进行二次加工FakeDataset用generate_data()批量生成随机Data对象后self.collate(data_list)非常适合在无法联网时快速验证模型与训练流程。其中download_url的实现见 torch_geometric/data/download.py它会把 URL 末尾的文件名作为保存名可通过filename参数覆盖若文件已存在则直接复用并打印Using existing file ...下载过程按 10MB 分块写入并同时暴露了download_google_url按 Google Drive 文件 ID 下载。此外 torch_geometric/data/init.py 还提供了extract_tar/extract_zip/extract_bz2/extract_gz等解压工具download方法里解压原始压缩包时可以按需调用。创建大规模数据集Dataset 与按需加载当数据无法整体装入内存时应使用Dataset基类。它紧密沿袭torchvision数据集的概念除了上述四个方法外还需要额外实现两个方法len()返回数据集中样本的数量get(idx)实现加载单个图的逻辑。其抽象签名见 dataset.py。内部机制上Dataset.__getitem__(idx)dataset.py会先调用self.get(self.indices()[idx])取回数据对象再按需应用transform若传入切片、列表、torch.Tensor/np.ndarraylong 或 bool 类型等索引则会走index_select返回数据集的子集视图。因此你只需保证get(idx)足够高效例如直接torch.load单文件__getitem__、__iter__、len、shuffle、index_select等派生能力便可免费获得。教程给出的Dataset实现示例如下——每个图数据对象在process中单独保存为data_{idx}.pt并在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注意该例中pre_filter采用不满足条件就continue跳过保存的方式pre_transform在保存前原地改写dataprocessed_file_names返回的列表长度即最终样本数因此len()直接取它的长度即可。此外若你的process会生成新文件还可以利用raw_paths/processed_paths属性dataset.py拿到raw_dir/processed_dir下所有文件的绝对路径列表与raw_file_names/processed_file_names一一对应。常见问题FAQ如何跳过download和/或process的执行只需不覆写download与process方法即可——has_download/has_process会返回False构造时便不会触发对应流程。教程给出的写法是class MyOwnDataset(Dataset): def __init__(self, transformNone, pre_transformNone): super().__init__(None, transform, pre_transform)这里rootNone表示不涉及磁盘读写完全在内存中工作Dataset.__init__会把root规范化为占位符MISSING见 dataset.py。如果只想在已处理数据上做随机变换甚至不需要继承数据集类。一定要使用这两套数据集接口吗不必须。正如在原生 PyTorch 中一样如果你想在飞行中生成合成数据、且不需要显式落盘可以直接把存放Data对象的普通 Python 列表交给DataLoaderfrom torch_geometric.data import Data from torch_geometric.loader import DataLoader data_list [Data(...), ..., Data(...)] loader DataLoader(data_list, batch_size32)torch_geometric/loader 中的DataLoader会基于collate自动把一批图拼成Batch对象这正是快速原型验证的捷径。另外Dataset还内置了get_summary()/print_summary()dataset.py统计数据集概览、to_datapipe()dataset.py转换为torch.utils.data.DataPipe、以及shuffle(return_perm...)等便捷方法值得在自定义数据集中直接继承复用。练习与解答考虑下面这个由Data对象列表构造的InMemoryDataset对应教程 docs/source/tutorial/create_dataset.rst 的 Exercises 一节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])问题 1self.processed_paths[0]的输出是什么答案是root目录下processed子目录中名为data.pt的绝对路径即os.path.join(self.root, processed, data.pt)。依据是 dataset.py 的processed_paths实现它把processed_file_names的结果此处为字符串data.pt会被to_list包装成[data.pt]逐一与processed_dir拼接。问题 2InMemoryDataset.save做了什么save是 PyG ≥ 2.4 引入的类方法它首先调用cls.collate(data_list)把样本列表拼接为单个Data对象并产出slices字典然后把三元组(data.to_dict(), slices, data.__class__)通过fs.torch_save写入指定路径in_memory_dataset.py。与之配套的load则读取该文件并按需恢复出self.data通过data_cls.from_dict重建对象与self.slices从而在__init__中完成数据集的内存化加载。若数据集只有一个样本collate会直接返回该样本、slices为None见 in_memory_dataset.py此时len()为 1。深入验证测试与更多扩展仓库的测试代码进一步印证了上述机制的预期行为test/data/test_dataset.py 覆盖了Dataset的len/get/index_select/shuffle以及切片索引行为test/data/test_inherit.py 验证了raw_file_names等属性被误写为普通方法时的兼容处理dataset.py 中isinstance(files, Callable)的防御逻辑test/data/test_data.py 与 test/loader/test_dataloader.py 则覆盖了Data对象与DataLoader的批量拼接链路。当内存仍显紧张时还可以调用InMemoryDataset.to_on_disk_dataset()in_memory_dataset.py把内存数据集一键转换为基于 SQLite 等后端的 OnDiskDataset适用于分布式训练或共享内存受限的硬件环境——这是教程之外 PyG 为超大数据集提供的另一条进阶路径。至此从目录约定、三个钩子函数到InMemoryDataset的四个方法与save/load/collate底层原理再到Dataset的按需加载与常见问题你已经掌握在 PyG 中构建自定义图数据集的完整方法论。动手实现你自己的MyOwnDataset配合DataLoader即可无缝接入下游的 GNN 训练与评测流程。【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表