Pytorch Geometric:如何在Colab中基于PyG Data对象列表创建临时小数据集
实现方案
PyG 没有官方提供类似 TensorDataset 的极简封装,但可以通过两种轻量方案实现内存级临时数据集,完全不需要持久化存储到磁盘:
方案1:兼容 PyG 全接口的最优方案(推荐)
直接继承 PyG 内置的 InMemoryDataset,跳过官方文档要求的下载、预处理逻辑,仅用你的Data列表完成初始化:
from torch_geometric.data import InMemoryDataset, Data from typing import List class TempPyGDataset(InMemoryDataset): def __init__(self, data_list: List[Data]): # 不需要传入root、transform等参数,直接初始化 super().__init__() # 直接调用PyG内置的collate方法处理数据列表 self.data, self.slices = self.collate(data_list) # 假设你已有的Data对象列表为 my_data_list dataset = TempPyGDataset(my_data_list)
该方案生成的数据集完全兼容 PyG 所有原生功能,包括自动获取num_node_features、num_classes等属性,以及配合DataLoader、随机划分等操作:
from torch.utils.data import random_split from torch_geometric.loader import DataLoader # 随机划分数据集 train_set, val_set, test_set = random_split(dataset, [0.8, 0.1, 0.1]) # 加载数据 train_loader = DataLoader(train_set, batch_size=32, shuffle=True)
方案2:极简自定义封装
如果你不需要用到 PyG 数据集的额外内置属性,完全可以对标 PyTorch 的自定义逻辑写最小实现,和你之前用TensorDataset的体验一致:
from torch.utils.data import Dataset from torch_geometric.data import Data from typing import List class SimplePyGDataset(Dataset): def __init__(self, data_list: List[Data]): self.data_list = data_list def __len__(self): return len(self.data_list) def __getitem__(self, idx): return self.data_list[idx] # 初始化方法和方案1完全一致 dataset = SimplePyGDataset(my_data_list)
该方案生成的数据集同样支持随机划分、PyG DataLoader加载等所有常规操作,仅缺少num_node_features这类PyG内置的数据集属性,需要的话手动给类加对应属性即可。
以上两种方案所有数据都仅保存在内存中,不会写入Colab磁盘,完全符合临时使用的需求,不需要参考官方文档里面向持久化存储的复杂配置流程。
内容的提问来源于stack exchange,提问作者CR-97
相关产品推荐
相关产品推荐

