InMemoryDataset 完整指南:3 步解决 PyTorch Geometric 图数据加载 OOM 问题
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
你刚把 PyTorch Geometric 的图数据加载逻辑写完,实例化数据集的那一刻内存直接爆掉,进程被系统杀掉。原因往往不是内存不够,而是存储策略和数据规模不匹配:PyG 内置了Dataset、InMemoryDataset、OnDiskDataset三种数据集基类,取舍点各不相同。下面以InMemoryDataset为主线,拆清楚它的合并存储机制、内存账本,以及内存吃紧时的转储路径。
InMemoryDataset、Dataset、OnDiskDataset 按规模怎么选 💡
先给判断表,30 秒对号入座:
| 基类 | 适用规模 | 存储形态 | 读取行为 | 典型场景 |
|---|---|---|---|---|
Dataset | 不限,取决于你自己实现get() | 每张图是独立对象/独立文件 | 按需读取,不合并 | 数据散落在多个文件、想逐图管理 |
InMemoryDataset | 整个数据集能装进 CPU 内存(建议占可用内存 2/3 以内) | 合并成单个Data+slices字典 | 一次性载入,单图读取走缓存 | Cora、QM9 这类中小规模数据集 |
OnDiskDataset | 超出 CPU 内存,或多进程要共读 | sqlite / rocksdb 数据库后端 | 磁盘随机读,不占大内存 | 大规模数据集、分布式训练 |
一句话口径:
- 放得进内存、单机单进程训练 →
InMemoryDataset; - 数据天然分散、或需要按图追加/更新 →
Dataset; - 装不进内存、或者多个训练进程要读同一份数据 →
OnDiskDataset。
InMemoryDataset 内存占用怎么算:合并存储 + 切片索引
InMemoryDataset不存 N 个图对象,而是把每张图的同名字段拼成一个大张量,这个拼接动作叫 collate(白话:把所有图的特征首尾接起来放一块)。合并的同时产出一个slices字典,记录每个样本在每个字段里的"起始位置和长度",作用类似书的目录。之后取任何一张图,就是照着目录从大张量里切出对应区间,实现都在 torch_geometric/data/in_memory_dataset.py:
# save:把 N 张图合并为 1 个 Data,slices 记录各字段切分位置 data, slices = self.collate(data_list) self.save((data, slices), self.processed_paths[0]) # load:构造函数一次性读回 self.load(self.processed_paths[0]) # get(idx):用 slices 从合并对象中切出第 idx 张图 # 之前取过的话,直接命中 _data_list 缓存并返回副本内存账本要算三层,很多人只算了第一层:
- 合并副本:约等于所有图张量体积之和,构造时
load()一次全部载入,这是基线; - 缓存副本:每次
get(idx)都会先用separate()切出一张单图,并在_data_list里留一份拷贝。跑完一个 epoch 后,手里同时握着原始合并数据 + N 份单图拷贝,体积约等于原始数据的 2 倍; - 批次副本:DataLoader 把每个 batch 的图再拼成一张大图训练,还会有一笔临时峰值。
所以估算公式是:预留内存 ≥ 2 × 数据集张量体积 + 1 个 batch 的体积。贴着上限跑的加载器,实际上早该 OOM 了,只是还没跑满一个 epoch 而已。
如何快速写出第一个 InMemoryDataset
第 1 步:自定义数据集的四件套模板
你只需要实现四处:两个*_file_names属性(告诉框架哪些文件齐全就跳过下载/处理),以及process()(把原始数据变成图列表并合并保存):
class MyOwnDataset(InMemoryDataset): def __init__(self, root, transform=None, pre_transform=None): super().__init__(root, transform, pre_transform) self.load(self.processed_paths[0]) @property def raw_file_names(self): return ['graphs.txt'] @property def processed_file_names(self): return ['data.pt'] def process(self): data, slices = self.collate([self._parse(i) for i in range(100)]) self.save((data, slices), self.processed_paths[0])为什么这样写:process()只在首次实例化时执行,之后构造器直接读processed目录的缓存文件。所以归一化、邻接矩阵求逆这类重计算要放进pre_transform(只跑一次,结果落盘);transform是每次访问都跑的,留给数据增强这类轻量操作。
第 2 步:图数据集接入 DataLoader
数据集写完直接交给 PyG 的DataLoader,batch 合并逻辑不用自己写——loader 对dataset[i]取出的单图自动拼接成大图:
from torch_geometric.loader import DataLoader loader = DataLoader(dataset, batch_size=32, shuffle=True) train_set = dataset.copy(dataset.train_idx) # 按索引切分子集 test_set = dataset.copy(dataset.test_idx)copy(idx)值得单独记住:传入切片、列表或索引张量都行,它会取出对应子集并重建合并存储,训练/验证/测试划分一行搞定,不用自己写拆分循环。
第 3 步:如何把图数据集转到磁盘 📦
内存吃紧时不用重写数据集,一行整体转成数据库文件:
on_disk = dataset.to_on_disk_dataset(root='data/on_disk', backend='sqlite') graph = on_disk[0] # 从磁盘按需读取它内部先从第一张图解析各字段类型生成 schema,再按每 1000 张图一批写入磁盘,读时按索引随机取,实现见 torch_geometric/data/on_disk_dataset.py。后端支持sqlite和rocksdb,数据量大、读写频繁时rocksdb更合适。
InMemoryDataset 解决不了什么:超大规模图与分布式训练的替代方案
⚠️ 合并存储有三个硬边界,碰到就该换方案:
- 超出 CPU 内存:合并后的张量必须一次性全部驻留内存。数据集到了 OGB-products 这个量级,只能走
OnDiskDataset路线——已有InMemoryDataset就直接to_on_disk_dataset()转储;数据还没进内存的话,继承OnDiskDataset重写process(),用append/extend分批写入即可。 - 分布式训练:合并对象只活在单个进程里,多机多卡没法共读。官方路径是先用 METIS 把大图切分,每台机器只持有本分区的拓扑和特征,再用
DistNeighborLoader做跨机邻居采样:代码入口在 examples/distributed/ 目录,从
pyg/下的partition_graph.py切分脚本跑起。 - 异构图转储受限:
to_on_disk_dataset()目前只支持同构Data,HeteroData需要自己继承OnDiskDataset,实现serialize/deserialize手动落库。
还有一个小坑:直接访问dataset.data会收到警告——那个合并对象是内部存储格式,字段请用dataset.{attr_name}访问,单图请用get(idx)。
延伸阅读:官方文档、示例与性能基准
- 数据集创建教程:docs/source/tutorial/create_dataset.rst
- 分布式训练教程:docs/source/tutorial/distributed_pyg.rst
- 分布式切分与采样示例:examples/distributed/
- 数据加载性能基准工具:benchmark/loader/
- 关键源码:torch_geometric/data/in_memory_dataset.py、torch_geometric/data/on_disk_dataset.py
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考