news 2026/9/5 22:24:13

InMemoryDataset 完整指南:3 步解决 PyTorch Geometric 图数据加载 OOM 问题

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
InMemoryDataset 完整指南:3 步解决 PyTorch Geometric 图数据加载 OOM 问题

InMemoryDataset 完整指南:3 步解决 PyTorch Geometric 图数据加载 OOM 问题

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

你刚把 PyTorch Geometric 的图数据加载逻辑写完,实例化数据集的那一刻内存直接爆掉,进程被系统杀掉。原因往往不是内存不够,而是存储策略和数据规模不匹配:PyG 内置了DatasetInMemoryDatasetOnDiskDataset三种数据集基类,取舍点各不相同。下面以InMemoryDataset为主线,拆清楚它的合并存储机制、内存账本,以及内存吃紧时的转储路径。

InMemoryDataset、Dataset、OnDiskDataset 按规模怎么选 💡

先给判断表,30 秒对号入座:

基类适用规模存储形态读取行为典型场景
Dataset不限,取决于你自己实现get()每张图是独立对象/独立文件按需读取,不合并数据散落在多个文件、想逐图管理
InMemoryDataset整个数据集能装进 CPU 内存(建议占可用内存 2/3 以内)合并成单个Data+slices字典一次性载入,单图读取走缓存Cora、QM9 这类中小规模数据集
OnDiskDataset超出 CPU 内存,或多进程要共读sqlite / rocksdb 数据库后端磁盘随机读,不占大内存大规模数据集、分布式训练

一句话口径:

  1. 放得进内存、单机单进程训练 →InMemoryDataset
  2. 数据天然分散、或需要按图追加/更新 →Dataset
  3. 装不进内存、或者多个训练进程要读同一份数据 →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 缓存并返回副本

内存账本要算三层,很多人只算了第一层:

  1. 合并副本:约等于所有图张量体积之和,构造时load()一次全部载入,这是基线;
  2. 缓存副本:每次get(idx)都会先用separate()切出一张单图,并在_data_list里留一份拷贝。跑完一个 epoch 后,手里同时握着原始合并数据 + N 份单图拷贝,体积约等于原始数据的 2 倍;
  3. 批次副本: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。后端支持sqliterocksdb,数据量大、读写频繁时rocksdb更合适。

InMemoryDataset 解决不了什么:超大规模图与分布式训练的替代方案

⚠️ 合并存储有三个硬边界,碰到就该换方案:

  1. 超出 CPU 内存:合并后的张量必须一次性全部驻留内存。数据集到了 OGB-products 这个量级,只能走OnDiskDataset路线——已有InMemoryDataset就直接to_on_disk_dataset()转储;数据还没进内存的话,继承OnDiskDataset重写process(),用append/extend分批写入即可。
  2. 分布式训练:合并对象只活在单个进程里,多机多卡没法共读。官方路径是先用 METIS 把大图切分,每台机器只持有本分区的拓扑和特征,再用DistNeighborLoader做跨机邻居采样:代码入口在 examples/distributed/ 目录,从pyg/下的partition_graph.py切分脚本跑起。
  3. 异构图转储受限to_on_disk_dataset()目前只支持同构DataHeteroData需要自己继承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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/5 22:21:00

Ice 菜单栏管理工具:macOS 图标收纳与自定义完整指南

Ice 菜单栏管理工具:macOS 图标收纳与自定义完整指南 【免费下载链接】Ice Powerful menu bar manager for macOS 项目地址: https://gitcode.com/GitHub_Trending/ice/Ice Ice 是一款运行在 macOS 上的菜单栏管理工具,核心任务是把你拥挤的菜单栏…

作者头像 李华
网站建设 2026/9/5 22:06:05

AI防幻觉基建:从RAG到本地模型的多层开源架构解析

先抛出今天这篇文章的核心观点:AI 之所以“不瞎编了”,不是因为模型突然变聪明了,而是因为工程上给它加了一圈“必须查资料、必须走流程、不允许自由发挥”的护栏。 这圈护栏并不是某一个框架能独立完成的。它在真实落地中往往由多层开源基建…

作者头像 李华
网站建设 2026/9/5 22:05:07

旅游知识图谱推荐系统:Python+Neo4j高分实战

简介:这是一套面向计算机专业本科生及研究生的高分课程实践资源,聚焦知识图谱技术在旅游推荐场景中的落地应用,适用于毕业设计、期末大作业与课程设计等中等难度实战项目。资源包含18个文件,以13个Python源码文件为核心&#xff0…

作者头像 李华