PyG DataLoader设置batch_size后无法正常生成多批次问题
问题原因
核心错误是导入了PyTorch原生的torch.utils.data.DataLoader,而非PyTorch Geometric(PyG)配套的图专用DataLoader。
- 原生PyTorch DataLoader的默认拼接函数(collate_fn)仅支持常规张量、数组类样本的分批拼接,无法识别PyG定义的
Data图结构对象,不会按照设定的batch_size参数拆分样本,只会将遍历到的所有样本一次性合并为单个大图,因此无论数据集内有多少个图文件,最终只会生成1个批次。 - 你观察到的输出中
ptr=[11]对应10个图样本、ptr=[2]对应1个图样本,就是原生loader错误地将全量样本直接拼接的典型表现,和你设置的batch_size=128参数完全无关。 - 你自定义的
GraphDataset类本身逻辑没有问题,__len__、__getitem__方法的实现均符合PyG数据集的规范要求,无需修改类代码。
修复方案
- 替换DataLoader的导入源,使用PyG官方提供的图数据专用加载器,注意不要和PyTorch原生DataLoader混用:
# 删除原有导入:from torch.utils.data import DataLoader from torch_geometric.loader import DataLoader - 保持原有的DataLoader实例化代码不变即可:
train_loader = DataLoader(train_set, batch_size=128) - 逻辑校验:替换后可通过
print(len(train_loader))确认分批逻辑是否正常:- 若训练集共10个图样本,由于10 < 128,返回1个批次属于正常表现;
- 若训练集样本量大于
batch_size(例如200个图文件),会自动返回对应数量的批次(200/128≈2个批次),不会再出现全量样本被合并为单个批次的问题。
内容的提问来源于stack exchange,提问作者alejandro maza villalpando
相关产品推荐
相关产品推荐

