PyTorch中如何合并两个torchvision.dataset对象为单个DataLoader
PyTorch合并CIFAR-10训练/测试集为单个DataLoader方案
核心思路是使用PyTorch官方提供的ConcatDataset工具拼接两个数据集对象,仅创建一个DataLoader即可解决你遇到的多进程超额、多次初始化开销问题:
import torchvision import torch from torch.utils.data import ConcatDataset, DataLoader batch_size = 128 # 加载原有两个数据集 cifar_trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=False) cifar_testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=False) # 拼接为全量数据集,无额外数据复制开销 full_cifar_set = ConcatDataset([cifar_trainset, cifar_testset]) # 仅初始化一个DataLoader,参数保持最优配置即可 cifar_full_dl = DataLoader( full_cifar_set, batch_size=batch_size, num_workers=12, persistent_workers=True, shuffle=True, pin_memory=True )
训练循环直接遍历单个DataLoader即可:
for data in cifar_full_dl: # 执行你的GAN训练逻辑 pass
该方案的优势:
- 总工作进程数保持为最优的12,不会出现进程超额抢占资源的问题
- 仅触发一次DataLoader初始化,没有多次启动的额外开销
- 开启
shuffle=True时会对全量60000张CIFAR-10样本全局打乱,比原方案分两次遍历数据集的方式更适合GAN训练,不会出现阶段性数据分布偏移的问题
注:你原有代码中的
cifar_l2为笔误,正确名称是cifar_dl2,使用本方案可直接规避该类错误。
内容的提问来源于stack exchange,提问作者IntegrateThis
相关产品推荐
相关产品推荐

